[llvm] [Semilattice] Introduce for dataflow analysis with KnownBits (PR #177616)

Ramkumar Ramachandra via llvm-commits llvm-commits at lists.llvm.org
Fri Jan 23 09:02:08 PST 2026


https://github.com/artagnon created https://github.com/llvm/llvm-project/pull/177616

Introduce a semilattice data structure holding KnownBits at each node. This will allow us to cache and invalidate KnownBits for a Value, invalidating only dataflow-dependent KnownBits.

The plan is to introduce a KnownBitsAnalysis, which computes the KnownBits for an entire function, keeping it in this semilattice. Eventually, we can migrate existing callers of computeKnownBits to use KnownBitsAnalysis, and ultimately remove the depth limitation on the function, leading to more complete results and lower compile times. It should be noted that, in the computation of KnownBits, we currently create copies of SimplifyQuery, changing certain things like the context-instruction, which would not work in this new world: hopefully, the regressions will be minor, and be paid for in lower compile times.

Testing was assisted by Claude AI.

-- 8< --
I've sketched a full working example in https://github.com/llvm/llvm-project/pull/176607 -- that would be the first intended use-case.

>From 6225bf5b45b635b0cffdea02bb4507db170c1ac4 Mon Sep 17 00:00:00 2001
From: Ramkumar Ramachandra <artagnon at tenstorrent.com>
Date: Tue, 20 Jan 2026 02:57:33 +0000
Subject: [PATCH] [Semilattice] Introduce for dataflow analysis with KnownBits

Introduce a semilattice data structure holding KnownBits at each node.
This will allow us to cache and invalidate KnownBits for a Value,
invalidating only dataflow-dependent KnownBits.

The plan is to introduce a KnownBitsAnalysis, which computes the
KnownBits for an entire function, keeping it in this semilattice.
Eventually, we can migrate existing callers of computeKnownBits to use
KnownBitsAnalysis, and ultimately remove the depth limitation on the
function, leading to more complete results and lower compile times. It
should be noted that, in the computation of KnownBits, we currently
create copies of SimplifyQuery, changing certain things like the
context-instruction, which would not work in this new world: hopefully,
the regressions will be minor, and be paid for in lower compile times.

Testing was assisted by Claude AI.
---
 llvm/include/llvm/Analysis/Semilattice.h    | 184 +++++++
 llvm/lib/Analysis/CMakeLists.txt            |   1 +
 llvm/lib/Analysis/Semilattice.cpp           | 108 +++++
 llvm/unittests/Analysis/CMakeLists.txt      |   1 +
 llvm/unittests/Analysis/SemilatticeTest.cpp | 506 ++++++++++++++++++++
 5 files changed, 800 insertions(+)
 create mode 100644 llvm/include/llvm/Analysis/Semilattice.h
 create mode 100644 llvm/lib/Analysis/Semilattice.cpp
 create mode 100644 llvm/unittests/Analysis/SemilatticeTest.cpp

diff --git a/llvm/include/llvm/Analysis/Semilattice.h b/llvm/include/llvm/Analysis/Semilattice.h
new file mode 100644
index 0000000000000..0079bca3423d6
--- /dev/null
+++ b/llvm/include/llvm/Analysis/Semilattice.h
@@ -0,0 +1,184 @@
+//===- Semilattice.h - A semilattice of KnownBits -------------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+// Builds a semilattice structure from the integral values in a given function,
+// and holds an associated KnownBits for each value, representing the dataflow
+// of KnownBits. Intended to be used to cache KnownBits for the entire function,
+// with invalidation APIs.
+//===----------------------------------------------------------------------===//
+
+#ifndef LLVM_ANALYSIS_SEMILATTICE_H
+#define LLVM_ANALYSIS_SEMILATTICE_H
+
+#include "llvm/ADT/DenseMap.h"
+#include "llvm/ADT/DepthFirstIterator.h"
+#include "llvm/ADT/GraphTraits.h"
+#include "llvm/ADT/PointerIntPair.h"
+#include "llvm/ADT/SmallVector.h"
+#include "llvm/IR/DerivedTypes.h"
+#include "llvm/IR/Value.h"
+#include "llvm/Support/Allocator.h"
+#include "llvm/Support/Compiler.h"
+#include "llvm/Support/KnownBits.h"
+
+namespace llvm {
+class Semilattice;
+
+class SemilatticeNode {
+  friend class Semilattice;
+  PointerIntPair<Value *, 1> ValHasKnownBits;
+  KnownBits Known;
+  SmallVector<SemilatticeNode *, 4> Parents;
+  SmallVector<SemilatticeNode *, 4> Children;
+
+public:
+  using iterator = SmallVectorImpl<SemilatticeNode *>::iterator;
+  LLVM_ABI_FOR_TEST iterator child_begin() { // NOLINT
+    return Children.begin();
+  }
+  LLVM_ABI_FOR_TEST iterator child_end() { return Children.end(); } // NOLINT
+  LLVM_ABI_FOR_TEST iterator_range<iterator> children() {
+    return make_range(child_begin(), child_end());
+  }
+  LLVM_ABI_FOR_TEST bool isLeaf() const { return Children.empty(); }
+  LLVM_ABI_FOR_TEST bool isRoot() const { return Parents.empty(); }
+  LLVM_ABI_FOR_TEST Value *getValue() const {
+    return ValHasKnownBits.getPointer();
+  }
+  LLVM_ABI_FOR_TEST bool hasKnownBits() const {
+    return ValHasKnownBits.getInt();
+  }
+  LLVM_ABI_FOR_TEST void setHasKnownBits() { ValHasKnownBits.setInt(1); }
+  LLVM_ABI_FOR_TEST KnownBits getKnownBits() const { return Known; }
+  LLVM_ABI_FOR_TEST void unionKnownWith(const KnownBits &NewKnown) {
+    Known = Known.unionWith(NewKnown);
+  }
+
+protected:
+  SemilatticeNode() : ValHasKnownBits(nullptr, 0) {}
+  SemilatticeNode(Value *V)
+      : ValHasKnownBits(V, 0),
+        Known(V->getType()->getScalarType()->getIntegerBitWidth()) {
+    assert(V && "Attempting to create a node with an empty Value");
+  }
+  SemilatticeNode(const SemilatticeNode &) = delete;
+  SemilatticeNode &operator=(const SemilatticeNode &) = delete;
+
+  void resetKnownBits() { Known.resetAll(); }
+  SemilatticeNode *addParent(SemilatticeNode *N) {
+    if (!hasParent(N))
+      Parents.push_back(N);
+    return this;
+  }
+  SemilatticeNode *addChild(SemilatticeNode *N) {
+    if (!is_contained(Children, N))
+      Children.push_back(N);
+    return this;
+  }
+  SemilatticeNode *rauw(Value *NewV) {
+    ValHasKnownBits.setPointer(NewV);
+    return this;
+  }
+  bool hasParent(SemilatticeNode *N) const { return is_contained(Parents, N); }
+};
+
+class Semilattice {
+  using NodeT = SemilatticeNode;
+
+  static constexpr size_t SlabSize = 8 * sizeof(NodeT);
+  BumpPtrAllocatorImpl<MallocAllocator, SlabSize, /*SizeThreshold=*/SlabSize,
+                       /*GrowthDelay=*/2>
+      NodeAllocator;
+
+  // The RootNode is a sentinel value to allow for graph traverals to work
+  // smoothly. Typically, to traverse the entire semilattice, a
+  // drop_begin(depth_first(Lat->getRootNode())) is used.
+  NodeT *RootNode;
+  DenseMap<Value *, NodeT *> NodeMap;
+  NodeT *create() { return new (NodeAllocator) NodeT(); }
+  NodeT *create(Value *V) { return new (NodeAllocator) NodeT(V); }
+  NodeT *getOrCreate(Value *V) { return NodeMap.lookup_or(V, create(V)); }
+  NodeT *insert(Value *V, NodeT *Parent = nullptr);
+  SmallVector<NodeT *, 4> insert_range(NodeT *Parent, // NOLINT
+                                       ArrayRef<User *> R);
+  void recurseInsertChildren(ArrayRef<NodeT *>);
+
+  // The roots (excluding the sentinel value) are the arguments of the function,
+  // and PHI nodes in each Basic Block, excluding values whose types are not
+  // either integer or vector of integers.
+  void initialize(ArrayRef<Value *> Roots);
+  void initialize(Function &F);
+
+public:
+  LLVM_ABI_FOR_TEST Semilattice(Function &F);
+  LLVM_ABI_FOR_TEST Semilattice(const Semilattice &) = delete;
+  LLVM_ABI_FOR_TEST Semilattice &operator=(const Semilattice &) = delete;
+  LLVM_ABI_FOR_TEST Semilattice(Semilattice &&) = default;
+  LLVM_ABI_FOR_TEST Semilattice &operator=(Semilattice &&) = default;
+
+  LLVM_ABI_FOR_TEST NodeT *getRootNode() const { return RootNode; }
+  LLVM_ABI_FOR_TEST bool empty() const { return RootNode->isLeaf(); }
+  LLVM_ABI_FOR_TEST bool contains(const Value *V) const {
+    return NodeMap.contains(V);
+  }
+  LLVM_ABI_FOR_TEST NodeT *lookup(const Value *V) const {
+    return NodeMap.lookup_or(V, RootNode);
+  }
+  LLVM_ABI_FOR_TEST size_t size() const { return NodeMap.size(); }
+
+  LLVM_ABI_FOR_TEST bool hasKnownBits(const Value *V) const {
+    return lookup(V)->hasKnownBits();
+  }
+  LLVM_ABI_FOR_TEST KnownBits getKnownBits(const Value *V) const {
+    return lookup(V)->getKnownBits();
+  }
+  LLVM_ABI_FOR_TEST void updateKnownBits(const Value *V,
+                                         const KnownBits &Known) const {
+    if (!contains(V))
+      return;
+    NodeT *LookupV = lookup(V);
+    LookupV->unionKnownWith(Known);
+    LookupV->setHasKnownBits();
+  }
+
+  // Theser functions return a reverse-breadth-first of the invalidated
+  // subgraph.
+  LLVM_ABI_FOR_TEST SmallVector<NodeT *> invalidateKnownBits(Value *V);
+  LLVM_ABI_FOR_TEST SmallVector<NodeT *> rauw(Value *OldV, Value *NewV);
+
+  LLVM_ABI_FOR_TEST void print(raw_ostream &OS) const;
+#if !defined(NDEBUG) || defined(LLVM_ENABLE_DUMP)
+  LLVM_DUMP_METHOD void dump() const;
+#endif
+};
+
+template <typename NodeRef> struct NodeGraphTraitsBase {
+  using ChildIteratorType = SemilatticeNode::iterator;
+  using nodes_iterator = df_iterator<NodeRef, df_iterator_default_set<NodeRef>>;
+
+  static NodeRef getEntryNode(NodeRef N) { return N; }
+  static ChildIteratorType child_begin(NodeRef N) { // NOLINT
+    return N->child_begin();
+  }
+  static ChildIteratorType child_end(NodeRef N) { // NOLINT
+    return N->child_end();
+  }
+};
+
+template <>
+struct GraphTraits<SemilatticeNode *>
+    : public NodeGraphTraitsBase<SemilatticeNode *> {
+  using NodeRef = SemilatticeNode *;
+};
+template <>
+struct GraphTraits<const SemilatticeNode *>
+    : public NodeGraphTraitsBase<const SemilatticeNode *> {
+  using NodeRef = const SemilatticeNode *;
+};
+} // end namespace llvm
+
+#endif // LLVM_ANALYSIS_SEMILATTICE_H
diff --git a/llvm/lib/Analysis/CMakeLists.txt b/llvm/lib/Analysis/CMakeLists.txt
index 9abdca099d9fe..15fabab57b8e2 100644
--- a/llvm/lib/Analysis/CMakeLists.txt
+++ b/llvm/lib/Analysis/CMakeLists.txt
@@ -141,6 +141,7 @@ add_llvm_component_library(LLVMAnalysis
   ScalarEvolutionAliasAnalysis.cpp
   ScalarEvolutionDivision.cpp
   ScalarEvolutionNormalization.cpp
+  Semilattice.cpp
   StaticDataProfileInfo.cpp
   StackLifetime.cpp
   StackSafetyAnalysis.cpp
diff --git a/llvm/lib/Analysis/Semilattice.cpp b/llvm/lib/Analysis/Semilattice.cpp
new file mode 100644
index 0000000000000..8fd6c275e526f
--- /dev/null
+++ b/llvm/lib/Analysis/Semilattice.cpp
@@ -0,0 +1,108 @@
+//===- Semilattice.cpp - A semilattice of KnownBits -----------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+
+#include "llvm/Analysis/Semilattice.h"
+#include "llvm/ADT/BreadthFirstIterator.h"
+#include "llvm/ADT/DepthFirstIterator.h"
+#include "llvm/ADT/STLExtras.h"
+#include "llvm/ADT/SetVector.h"
+#include "llvm/IR/Function.h"
+#include "llvm/Support/Debug.h"
+
+using namespace llvm;
+using NodeT = SemilatticeNode;
+
+void Semilattice::initialize(ArrayRef<Value *> Roots) {
+  auto ToInsert = make_filter_range(Roots, [&](Value *V) {
+    return !contains(V) && V->getType()->isIntOrIntVectorTy();
+  });
+  for (Value *V : ToInsert)
+    recurseInsertChildren(insert(V));
+}
+
+void Semilattice::initialize(Function &F) {
+  SmallVector<Value *> Args(llvm::make_pointer_range(F.args()));
+  initialize(Args);
+  for (BasicBlock &BB : F)
+    if (!BB.empty()) {
+      SmallVector<Value *> Args(llvm::make_pointer_range(
+          make_range(BB.begin(), BB.getFirstNonPHIIt())));
+      initialize(Args);
+    }
+}
+
+Semilattice::Semilattice(Function &F) : RootNode(create()) { initialize(F); }
+
+NodeT *Semilattice::insert(Value *V, NodeT *Parent) {
+  assert(V->getType()->isIntOrIntVectorTy() &&
+         "Cannot insert non-integral types");
+  NodeT *Node = getOrCreate(V);
+  NodeMap.try_emplace(V, Node);
+  NodeT *ParentNode = Parent ? Parent : RootNode;
+  ParentNode->addChild(Node);
+  return Node->addParent(ParentNode);
+}
+
+SmallVector<NodeT *, 4> Semilattice::insert_range(NodeT *Parent,
+                                                  ArrayRef<User *> R) {
+  SmallVector<NodeT *, 4> Ret;
+  auto Users = make_filter_range(R, [&](User *U) {
+    return (!contains(U) || !lookup(U)->hasParent(Parent)) &&
+           U->getType()->isIntOrIntVectorTy();
+  });
+  for (User *U : Users)
+    Ret.push_back(insert(U, Parent));
+  return Ret;
+}
+
+void Semilattice::recurseInsertChildren(ArrayRef<NodeT *> Parents) {
+  for (NodeT *P : Parents)
+    recurseInsertChildren(insert_range(P, to_vector(P->getValue()->users())));
+}
+
+SmallVector<NodeT *> Semilattice::invalidateKnownBits(Value *V) {
+  if (!contains(V))
+    return {};
+  SetVector<NodeT *> ToUpdate;
+  for (NodeT *N : breadth_first(lookup(V))) {
+    N->resetKnownBits();
+    N->setHasKnownBits();
+    ToUpdate.insert(N);
+  }
+  SmallVector<NodeT *> Ret(reverse(ToUpdate.takeVector()));
+  return Ret;
+}
+
+SmallVector<NodeT *> Semilattice::rauw(Value *OldV, Value *NewV) {
+  if (!contains(OldV))
+    return {};
+  NodeMap.emplace_or_assign(NewV, lookup(OldV)->rauw(NewV));
+  NodeMap.erase(OldV);
+  return invalidateKnownBits(NewV);
+}
+
+void Semilattice::print(raw_ostream &OS) const {
+  for (NodeT *N : drop_begin(depth_first(RootNode))) {
+    if (N->hasParent(RootNode))
+      OS << "^ ";
+    else if (N->isLeaf())
+      OS << "$ ";
+    else
+      OS << "  ";
+    N->getValue()->print(OS);
+    if (N->hasKnownBits()) {
+      OS << " | ";
+      N->Known.print(OS);
+    }
+    OS << "\n";
+  }
+}
+
+#if !defined(NDEBUG) || defined(LLVM_ENABLE_DUMP)
+LLVM_DUMP_METHOD void Semilattice::dump() const { print(dbgs()); }
+#endif
diff --git a/llvm/unittests/Analysis/CMakeLists.txt b/llvm/unittests/Analysis/CMakeLists.txt
index 50bf4539e7984..852ae5d449a9c 100644
--- a/llvm/unittests/Analysis/CMakeLists.txt
+++ b/llvm/unittests/Analysis/CMakeLists.txt
@@ -51,6 +51,7 @@ set(ANALYSIS_TEST_SOURCES
   ProfileSummaryInfoTest.cpp
   ReplaceWithVecLibTest.cpp
   ScalarEvolutionTest.cpp
+  SemilatticeTest.cpp
   SparsePropagation.cpp
   TargetLibraryInfoTest.cpp
   TensorSpecTest.cpp
diff --git a/llvm/unittests/Analysis/SemilatticeTest.cpp b/llvm/unittests/Analysis/SemilatticeTest.cpp
new file mode 100644
index 0000000000000..320cb5bd0f0e4
--- /dev/null
+++ b/llvm/unittests/Analysis/SemilatticeTest.cpp
@@ -0,0 +1,506 @@
+//===- SemilatticeTest.cpp - Semilattice tests ----------------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+
+#include "llvm/Analysis/Semilattice.h"
+#include "llvm/ADT/DepthFirstIterator.h"
+#include "llvm/ADT/STLExtras.h"
+#include "llvm/AsmParser/Parser.h"
+#include "llvm/IR/BasicBlock.h"
+#include "llvm/IR/Function.h"
+#include "llvm/IR/InstIterator.h"
+#include "llvm/IR/LLVMContext.h"
+#include "llvm/IR/Module.h"
+#include "llvm/Support/SourceMgr.h"
+#include "gtest/gtest.h"
+
+using namespace llvm;
+
+namespace {
+static Instruction *findInstructionByName(Function *F, StringRef Name) {
+  for (Instruction &I : instructions(F))
+    if (I.getName() == Name)
+      return &I;
+  return nullptr;
+}
+
+class SemilatticeTest : public ::testing::Test {
+protected:
+  void parseAssembly(StringRef Assembly) {
+    SMDiagnostic Error;
+    M = parseAssemblyString(Assembly, Error, Context);
+    ASSERT_TRUE(M);
+    F = M->getFunction("test");
+    ASSERT_TRUE(F) << "Test must have a function @test";
+    if (!F)
+      return;
+    Counter = findInstructionByName(F, "counter");
+    NextCounter = findInstructionByName(F, "next_counter");
+    Result = findInstructionByName(F, "result");
+    PhiCounter = findInstructionByName(F, "phi_counter");
+    PhiResult = findInstructionByName(F, "phi_result");
+    Cond = findInstructionByName(F, "cond");
+    BranchVal = findInstructionByName(F, "branch_val");
+    MergeVal = findInstructionByName(F, "merge_val");
+  }
+  void SetUp() override { M = std::make_unique<Module>("test", Context); }
+  Function *createSimpleFunction(StringRef Name,
+                                 ArrayRef<Type *> ArgTypes = {}) {
+    std::vector<Type *> Types(ArgTypes.begin(), ArgTypes.end());
+    FunctionType *FTy =
+        FunctionType::get(Type::getVoidTy(Context), Types, false);
+    return Function::Create(FTy, Function::ExternalLinkage, Name, M.get());
+  }
+  LLVMContext Context;
+  std::unique_ptr<Module> M;
+  Function *F = nullptr;
+  Instruction *Counter = nullptr, *NextCounter = nullptr, *Result = nullptr,
+              *PhiCounter = nullptr, *PhiResult = nullptr, *Cond = nullptr,
+              *BranchVal = nullptr, *MergeVal = nullptr;
+};
+
+TEST_F(SemilatticeTest, BasicConstruction) {
+  parseAssembly(
+      "define void @test(i32 %n) {\n"
+      "entry:\n"
+      "  br label %loop\n"
+      "loop:\n"
+      "  %phi_counter = phi i32 [ 0, %entry ], [ %next_counter, %loop ]\n"
+      "  %phi_result = phi i32 [ 1, %entry ], [ %result, %loop ]\n"
+      "  %counter = add i32 %phi_counter, 1\n"
+      "  %result = mul i32 %phi_result, 2\n"
+      "  %next_counter = add i32 %counter, 1\n"
+      "  %cond = icmp slt i32 %next_counter, %n\n"
+      "  br i1 %cond, label %loop, label %exit\n"
+      "exit:\n"
+      "  store i32 %result, ptr poison\n"
+      "  ret void\n"
+      "}\n");
+  Semilattice Lat(*F);
+  Argument *ArgN = &*F->arg_begin();
+  EXPECT_TRUE(Lat.contains(ArgN));
+  EXPECT_TRUE(Lat.contains(PhiCounter));
+  EXPECT_TRUE(Lat.contains(PhiResult));
+  EXPECT_TRUE(Lat.contains(Counter));
+  EXPECT_TRUE(Lat.contains(Result));
+  EXPECT_TRUE(Lat.contains(NextCounter));
+  EXPECT_TRUE(Lat.contains(Cond));
+  EXPECT_EQ(Lat.size(), 7u);
+  EXPECT_FALSE(Lat.lookup(PhiCounter)->isLeaf());
+  EXPECT_FALSE(Lat.lookup(PhiResult)->isLeaf());
+  EXPECT_FALSE(Lat.lookup(Counter)->isLeaf());
+  EXPECT_FALSE(Lat.lookup(Result)->isLeaf());
+  EXPECT_FALSE(Lat.lookup(NextCounter)->isLeaf());
+  EXPECT_TRUE(Lat.lookup(Cond)->isLeaf());
+}
+
+TEST_F(SemilatticeTest, ConstructionNonIntegralExcluded) {
+  parseAssembly(
+      "define void @test(i32 %int_arg, float %float_arg, ptr %ptr_arg, <2 x "
+      "i32> %vec_arg) {\n"
+      "entry:\n"
+      "  br i1 poison, label %then, label %else\n"
+      "then:\n"
+      "  %int_val = add i32 %int_arg, 1\n"
+      "  %float_val = fadd float %float_arg, 1.0\n"
+      "  %vec_val = add <2 x i32> %vec_arg, <i32 1, i32 2>\n"
+      "  br label %merge\n"
+      "else:\n"
+      "  %int_val2 = mul i32 %int_arg, 2\n"
+      "  %ptr_val = getelementptr i8, ptr %ptr_arg, i32 4\n"
+      "  %vec_val2 = mul <2 x i32> %vec_arg, <i32 3, i32 4>\n"
+      "  br label %merge\n"
+      "merge:\n"
+      "  %phi_int = phi i32 [ %int_val, %then ], [ %int_val2, %else ]\n"
+      "  %phi_float = phi float [ %float_val, %then ], [ %float_arg, %else ]\n"
+      "  %phi_ptr = phi ptr [ %ptr_arg, %then ], [ %ptr_val, %else ]\n"
+      "  %phi_vec = phi <2 x i32> [ %vec_val, %then ], [ %vec_val2, %else ]\n"
+      "  %final_int = add i32 %phi_int, 5\n"
+      "  %final_vec = add <2 x i32> %phi_vec, <i32 5, i32 6>\n"
+      "  store float %phi_float, ptr %phi_ptr\n"
+      "  ret void\n"
+      "}\n");
+  Semilattice Lat(*F);
+  auto *ArgIt = F->arg_begin();
+  Argument *IntArg = &*ArgIt++;
+  Argument *FloatArg = &*ArgIt++;
+  Argument *PtrArg = &*ArgIt++;
+  Argument *VecArg = &*ArgIt;
+  Instruction *IntVal = findInstructionByName(F, "int_val");
+  Instruction *FloatVal = findInstructionByName(F, "float_val");
+  Instruction *VecVal = findInstructionByName(F, "vec_val");
+  Instruction *IntVal2 = findInstructionByName(F, "int_val2");
+  Instruction *PtrVal = findInstructionByName(F, "ptr_val");
+  Instruction *VecVal2 = findInstructionByName(F, "vec_val2");
+  Instruction *PhiInt = findInstructionByName(F, "phi_int");
+  Instruction *PhiFloat = findInstructionByName(F, "phi_float");
+  Instruction *PhiPtr = findInstructionByName(F, "phi_ptr");
+  Instruction *PhiVec = findInstructionByName(F, "phi_vec");
+  Instruction *FinalInt = findInstructionByName(F, "final_int");
+  Instruction *FinalVec = findInstructionByName(F, "final_vec");
+
+  EXPECT_TRUE(Lat.contains(IntArg));
+  EXPECT_TRUE(Lat.contains(IntVal));
+  EXPECT_TRUE(Lat.contains(IntVal2));
+  EXPECT_TRUE(Lat.contains(PhiInt));
+  EXPECT_TRUE(Lat.contains(FinalInt));
+
+  EXPECT_TRUE(Lat.contains(VecArg));
+  EXPECT_TRUE(Lat.contains(VecVal));
+  EXPECT_TRUE(Lat.contains(VecVal2));
+  EXPECT_TRUE(Lat.contains(PhiVec));
+  EXPECT_TRUE(Lat.contains(FinalVec));
+
+  EXPECT_FALSE(Lat.contains(FloatArg));
+  EXPECT_FALSE(Lat.contains(PtrArg));
+  EXPECT_FALSE(Lat.contains(FloatVal));
+  EXPECT_FALSE(Lat.contains(PtrVal));
+  EXPECT_FALSE(Lat.contains(PhiFloat));
+  EXPECT_FALSE(Lat.contains(PhiPtr));
+
+  EXPECT_EQ(Lat.size(), 10u);
+}
+
+TEST_F(SemilatticeTest, Iteration) {
+  parseAssembly("define void @test(i32 %arg1, i32 %arg2) {\n"
+                "  %add1 = add i32 %arg1, 1\n"
+                "  %add2 = add i32 %arg2, 2\n"
+                "  %mul1 = mul i32 %add1, 3\n"
+                "  %mul2 = mul i32 %add2, 4\n"
+                "  %final = add i32 %mul1, %mul2\n"
+                "  ret void\n"
+                "}\n");
+  Semilattice Lat(*F);
+  auto *ArgIt = F->arg_begin();
+  Argument *Arg1 = &*ArgIt++;
+  Argument *Arg2 = &*ArgIt;
+  Instruction *Add1 = findInstructionByName(F, "add1");
+  Instruction *Mul1 = findInstructionByName(F, "mul1");
+  Instruction *Final = findInstructionByName(F, "final");
+
+  SemilatticeNode *RootNode = Lat.getRootNode();
+  EXPECT_FALSE(RootNode->isLeaf());
+  SmallVector<SemilatticeNode *> RootChildren(RootNode->children());
+  EXPECT_EQ(RootChildren.size(), 2u);
+  EXPECT_EQ(RootChildren[0]->getValue(), Arg1);
+  EXPECT_EQ(RootChildren[1]->getValue(), Arg2);
+
+  SemilatticeNode *Arg1Node = Lat.lookup(Arg1);
+  EXPECT_FALSE(Arg1Node->isLeaf());
+  SmallVector<SemilatticeNode *> Arg1Children(Arg1Node->children());
+  EXPECT_EQ(Arg1Children.size(), 1u);
+  EXPECT_EQ(Arg1Children[0]->getValue(), Add1);
+
+  SemilatticeNode *Add1Node = Lat.lookup(Add1);
+  EXPECT_FALSE(Add1Node->isLeaf());
+  SmallVector<SemilatticeNode *> Add1Children(Add1Node->children());
+  EXPECT_EQ(Add1Children.size(), 1u);
+  EXPECT_EQ(Add1Children[0]->getValue(), Mul1);
+
+  SemilatticeNode *Mul1Node = Lat.lookup(Mul1);
+  EXPECT_FALSE(Mul1Node->isLeaf());
+  SmallVector<SemilatticeNode *> Mul1Children(Mul1Node->children());
+  EXPECT_EQ(Mul1Children.size(), 1u);
+  EXPECT_EQ(Mul1Children[0]->getValue(), Final);
+
+  SemilatticeNode *FinalNode = Lat.lookup(Final);
+  EXPECT_TRUE(FinalNode->isLeaf());
+  SmallVector<SemilatticeNode *> FinalChildren(FinalNode->children());
+  EXPECT_EQ(FinalChildren.size(), 0u);
+
+  SmallVector<SemilatticeNode *> DepthFirstOrder(depth_first(RootNode));
+  EXPECT_GE(DepthFirstOrder.size(), 7u);
+  EXPECT_EQ(DepthFirstOrder[0], RootNode);
+
+  auto *RootPos = find(DepthFirstOrder, RootNode);
+  auto *Arg1NodePos = find(DepthFirstOrder, Arg1Node);
+  auto *Add1NodePos = find(DepthFirstOrder, Add1Node);
+  auto *Mul1NodePos = find(DepthFirstOrder, Mul1Node);
+  auto *FinalNodePos = find(DepthFirstOrder, FinalNode);
+
+  EXPECT_NE(RootPos, DepthFirstOrder.end());
+  EXPECT_NE(Arg1NodePos, DepthFirstOrder.end());
+  EXPECT_NE(Add1NodePos, DepthFirstOrder.end());
+  EXPECT_NE(Mul1NodePos, DepthFirstOrder.end());
+  EXPECT_NE(FinalNodePos, DepthFirstOrder.end());
+  EXPECT_LT(RootPos, Arg1NodePos);
+  EXPECT_LT(Arg1NodePos, Add1NodePos);
+  EXPECT_LT(Add1NodePos, Mul1NodePos);
+  EXPECT_LT(Mul1NodePos, FinalNodePos);
+}
+
+TEST_F(SemilatticeTest, NestedLoop) {
+  parseAssembly(
+      "define void @test(i32 %n, i32 %m) {\n"
+      "entry:\n"
+      "  br label %outer_loop\n"
+      "outer_loop:\n"
+      "  %outer_phi = phi i32 [ 0, %entry ], [ %outer_next, %outer_latch ]\n"
+      "  br label %inner_loop\n"
+      "inner_loop:\n"
+      "  %inner_phi = phi i32 [ 0, %outer_loop ], [ %inner_next, %inner_loop "
+      "]\n"
+      "  %inner_next = add i32 %inner_phi, 1\n"
+      "  %inner_cond = icmp slt i32 %inner_next, %m\n"
+      "  br i1 %inner_cond, label %inner_loop, label %outer_latch\n"
+      "outer_latch:\n"
+      "  %outer_next = add i32 %outer_phi, 1\n"
+      "  %outer_cond = icmp slt i32 %outer_next, %n\n"
+      "  br i1 %outer_cond, label %outer_loop, label %exit\n"
+      "exit:\n"
+      "  ret void\n"
+      "}\n");
+  Semilattice Lat(*F);
+  auto *ArgIt = F->arg_begin();
+  Argument *ArgN = &*ArgIt++;
+  Argument *ArgM = &*ArgIt;
+  Instruction *OuterPHI = findInstructionByName(F, "outer_phi");
+  Instruction *InnerPHI = findInstructionByName(F, "inner_phi");
+  Instruction *InnerNext = findInstructionByName(F, "inner_next");
+  Instruction *OuterNext = findInstructionByName(F, "outer_next");
+  Instruction *InnerCond = findInstructionByName(F, "inner_cond");
+  Instruction *OuterCond = findInstructionByName(F, "outer_cond");
+
+  EXPECT_TRUE(Lat.contains(ArgN));
+  EXPECT_TRUE(Lat.contains(ArgM));
+  EXPECT_TRUE(Lat.contains(OuterPHI));
+  EXPECT_TRUE(Lat.contains(InnerPHI));
+  EXPECT_TRUE(Lat.contains(InnerNext));
+  EXPECT_TRUE(Lat.contains(OuterNext));
+  EXPECT_TRUE(Lat.contains(InnerCond));
+  EXPECT_TRUE(Lat.contains(OuterCond));
+  EXPECT_EQ(Lat.size(), 8u);
+}
+
+TEST_F(SemilatticeTest, InvalidateKnownBitsSubgraph) {
+  parseAssembly("define void @test(i32 %arg) {\n"
+                "  %counter = add i32 %arg, 1\n"
+                "  %result = mul i32 %counter, 2\n"
+                "  %next_counter = add i32 %result, 3\n"
+                "  %branch_val = sub i32 %next_counter, 1\n"
+                "  %merge_val = add i32 %branch_val, 5\n"
+                "  store i32 %merge_val, ptr poison\n"
+                "  ret void\n"
+                "}\n");
+  Semilattice Lat(*F);
+  Argument *Arg = &*F->arg_begin();
+  KnownBits Known32(32);
+  Known32.setAllOnes();
+
+  Lat.updateKnownBits(Arg, Known32);
+  Lat.updateKnownBits(Counter, Known32);
+  Lat.updateKnownBits(Result, Known32);
+  Lat.updateKnownBits(NextCounter, Known32);
+  Lat.updateKnownBits(BranchVal, Known32);
+  Lat.updateKnownBits(MergeVal, Known32);
+
+  EXPECT_TRUE(Lat.getKnownBits(Arg).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(Counter).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(Result).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(NextCounter).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(BranchVal).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(MergeVal).isAllOnes());
+
+  SmallVector<SemilatticeNode *> InvalidatedNodes =
+      Lat.invalidateKnownBits(Counter);
+  EXPECT_TRUE(Lat.getKnownBits(Arg).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(Counter).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(Result).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(NextCounter).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(BranchVal).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(MergeVal).isUnknown());
+  EXPECT_EQ(InvalidatedNodes.size(), 5u);
+}
+
+TEST_F(SemilatticeTest, InvalidateKnownBitsPhiSubgraph) {
+  parseAssembly(
+      "define void @test(i32 %n, i1 %cond) {\n"
+      "entry:\n"
+      "  %counter = add i32 %n, 1\n"
+      "  br i1 %cond, label %then, label %else\n"
+      "then:\n"
+      "  %branch_val = mul i32 %counter, 2\n"
+      "  br label %merge\n"
+      "else:\n"
+      "  %result = add i32 %counter, 3\n"
+      "  br label %merge\n"
+      "merge:\n"
+      "  %merge_val = phi i32 [ %branch_val, %then ], [ %result, %else ]\n"
+      "  %next_counter = add i32 %merge_val, 1\n"
+      "  store i32 %next_counter, ptr poison\n"
+      "  ret void\n"
+      "}\n");
+  Semilattice Lat(*F);
+  auto *ArgIt = F->arg_begin();
+  Argument *ArgN = &*ArgIt++;
+  Argument *ArgCond = &*ArgIt;
+  KnownBits Known32(32);
+  Known32.setAllOnes();
+  KnownBits Known1(1);
+  Known1.setAllOnes();
+
+  Lat.updateKnownBits(ArgN, Known32);
+  Lat.updateKnownBits(ArgCond, Known1);
+  Lat.updateKnownBits(Counter, Known32);
+  Lat.updateKnownBits(BranchVal, Known32);
+  Lat.updateKnownBits(Result, Known32);
+  Lat.updateKnownBits(MergeVal, Known32);
+  Lat.updateKnownBits(NextCounter, Known32);
+
+  EXPECT_TRUE(Lat.getKnownBits(Counter).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(BranchVal).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(Result).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(MergeVal).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(NextCounter).isAllOnes());
+
+  SmallVector<SemilatticeNode *> InvalidatedNodes =
+      Lat.invalidateKnownBits(Counter);
+  EXPECT_TRUE(Lat.getKnownBits(ArgN).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(ArgCond).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(Counter).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(BranchVal).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(Result).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(MergeVal).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(NextCounter).isUnknown());
+  EXPECT_EQ(InvalidatedNodes.size(), 5u);
+}
+
+TEST_F(SemilatticeTest, RauwSubgraphInvalidation) {
+  parseAssembly("define void @test(i32 %arg1, i32 %arg2) {\n"
+                "  %counter = add i32 %arg1, 1\n"
+                "  %result = mul i32 %counter, 2\n"
+                "  %next_counter = add i32 %result, %arg2\n"
+                "  store i32 %next_counter, ptr poison\n"
+                "  ret void\n"
+                "}\n");
+  Semilattice Lat(*F);
+  Function *F2 =
+      createSimpleFunction("other_func", {Type::getInt32Ty(Context)});
+  Argument *NewArg = &*F2->arg_begin();
+  auto *ArgIt = F->arg_begin();
+  Argument *Arg1 = &*ArgIt++;
+  Argument *Arg2 = &*ArgIt;
+  KnownBits Known32(32);
+  Known32.setAllOnes();
+
+  Lat.updateKnownBits(Arg1, Known32);
+  Lat.updateKnownBits(Arg2, Known32);
+  Lat.updateKnownBits(Counter, Known32);
+  Lat.updateKnownBits(Result, Known32);
+  Lat.updateKnownBits(NextCounter, Known32);
+
+  EXPECT_TRUE(Lat.getKnownBits(Arg1).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(Counter).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(Result).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(NextCounter).isAllOnes());
+  EXPECT_TRUE(Lat.contains(Arg1));
+  EXPECT_FALSE(Lat.contains(NewArg));
+
+  SmallVector<SemilatticeNode *> InvalidatedNodes = Lat.rauw(Arg1, NewArg);
+  EXPECT_FALSE(Lat.contains(Arg1));
+  EXPECT_TRUE(Lat.contains(NewArg));
+  EXPECT_TRUE(Lat.getKnownBits(NewArg).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(Counter).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(Result).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(NextCounter).isUnknown());
+  EXPECT_EQ(InvalidatedNodes.size(), 4u);
+}
+
+TEST_F(SemilatticeTest, RauwPhiSubgraphInvalidation) {
+  parseAssembly(
+      "define void @test(i32 %n, i1 %cond) {\n"
+      "entry:\n"
+      "  br i1 %cond, label %then, label %else\n"
+      "then:\n"
+      "  %branch_val = add i32 %n, 5\n"
+      "  br label %merge\n"
+      "else:\n"
+      "  %result = mul i32 %n, 3\n"
+      "  br label %merge\n"
+      "merge:\n"
+      "  %merge_val = phi i32 [ %branch_val, %then ], [ %result, %else ]\n"
+      "  %counter = add i32 %merge_val, 1\n"
+      "  %next_counter = mul i32 %counter, 2\n"
+      "  store i32 %next_counter, ptr poison\n"
+      "  ret void\n"
+      "}\n");
+  Semilattice Lat(*F);
+  Function *F2 =
+      createSimpleFunction("other_func", {Type::getInt32Ty(Context)});
+  Argument *NewArg = &*F2->arg_begin();
+  auto *ArgIt = F->arg_begin();
+  Argument *ArgN = &*ArgIt++;
+  Argument *ArgCond = &*ArgIt;
+  KnownBits Known32(32);
+  Known32.setAllOnes();
+  KnownBits Known1(1);
+  Known1.setAllOnes();
+
+  Lat.updateKnownBits(ArgN, Known32);
+  Lat.updateKnownBits(ArgCond, Known1);
+  Lat.updateKnownBits(BranchVal, Known32);
+  Lat.updateKnownBits(Result, Known32);
+  Lat.updateKnownBits(MergeVal, Known32);
+  Lat.updateKnownBits(Counter, Known32);
+  Lat.updateKnownBits(NextCounter, Known32);
+
+  EXPECT_TRUE(Lat.getKnownBits(ArgN).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(ArgCond).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(BranchVal).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(Result).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(MergeVal).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(Counter).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(NextCounter).isAllOnes());
+  EXPECT_TRUE(Lat.contains(ArgN));
+  EXPECT_FALSE(Lat.contains(NewArg));
+
+  SmallVector<SemilatticeNode *> InvalidatedNodes = Lat.rauw(ArgN, NewArg);
+  EXPECT_FALSE(Lat.contains(ArgN));
+  EXPECT_TRUE(Lat.contains(NewArg));
+  EXPECT_TRUE(Lat.getKnownBits(ArgCond).isAllOnes());
+  EXPECT_TRUE(Lat.getKnownBits(NewArg).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(BranchVal).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(Result).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(MergeVal).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(Counter).isUnknown());
+  EXPECT_TRUE(Lat.getKnownBits(NextCounter).isUnknown());
+  EXPECT_EQ(InvalidatedNodes.size(), 6u);
+}
+
+TEST_F(SemilatticeTest, Print) {
+  parseAssembly(
+      "define void @test(i32 %n) {\n"
+      "entry:\n"
+      "  br label %loop\n"
+      "loop:\n"
+      "  %phi_counter = phi i32 [ 0, %entry ], [ %next_counter, %loop ]\n"
+      "  %counter = add i32 %phi_counter, 1\n"
+      "  %result = mul i32 %counter, 2\n"
+      "  %next_counter = add i32 %result, 1\n"
+      "  %cond = icmp slt i32 %next_counter, %n\n"
+      "  br i1 %cond, label %loop, label %exit\n"
+      "exit:\n"
+      "  ret void\n"
+      "}\n");
+  Semilattice Lat(*F);
+  KnownBits Known32(32);
+  Known32.setAllZero();
+  Instruction *Result = findInstructionByName(F, "result");
+  Lat.updateKnownBits(Result, Known32);
+  std::string ActualOutput;
+  raw_string_ostream OS(ActualOutput);
+  Lat.print(OS);
+  std::string ExpectedOutput =
+      "^ i32 %n\n"
+      "$   %cond = icmp slt i32 %next_counter, %n\n"
+      "^   %phi_counter = phi i32 [ 0, %entry ], [ %next_counter, %loop ]\n"
+      "    %counter = add i32 %phi_counter, 1\n"
+      "    %result = mul i32 %counter, 2 | 00000000000000000000000000000000\n"
+      "    %next_counter = add i32 %result, 1\n";
+  EXPECT_EQ(ActualOutput, ExpectedOutput);
+}
+} // namespace



More information about the llvm-commits mailing list