[llvm] [LAA] getPointersDiff(): Add ExpensivePtrCheck argument (PR #226333)

via llvm-commits llvm-commits at lists.llvm.org
Thu Sep 24 17:48:48 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-llvm-analysis

Author: Vasileios Porpodas (vporpo)

<details>
<summary>Changes</summary>

This patch adds the `ExpensivePtrCheck` argument to getPointersDiff(). When this is true we use SE.getMinusSCEV() instead of SE.computeConstantDifference(), which can return more accurate results if the pointer computation is complex.

---
Full diff: https://github.com/llvm/llvm-project/pull/226333.diff


4 Files Affected:

- (modified) llvm/include/llvm/Analysis/LoopAccessAnalysis.h (+2-1) 
- (modified) llvm/lib/Analysis/LoopAccessAnalysis.cpp (+15-3) 
- (modified) llvm/unittests/Analysis/CMakeLists.txt (+1) 
- (added) llvm/unittests/Analysis/LoopAccessAnalysisTest.cpp (+84) 


``````````diff
diff --git a/llvm/include/llvm/Analysis/LoopAccessAnalysis.h b/llvm/include/llvm/Analysis/LoopAccessAnalysis.h
index 792b8ba9cabc5..67bcb09b39d0b 100644
--- a/llvm/include/llvm/Analysis/LoopAccessAnalysis.h
+++ b/llvm/include/llvm/Analysis/LoopAccessAnalysis.h
@@ -961,7 +961,8 @@ getPtrStride(PredicatedScalarEvolution &PSE, Type *AccessTy, Value *Ptr,
 LLVM_ABI std::optional<int64_t>
 getPointersDiff(Type *ElemTyA, Value *PtrA, Type *ElemTyB, Value *PtrB,
                 const DataLayout &DL, ScalarEvolution &SE,
-                bool StrictCheck = false, bool CheckType = true);
+                bool StrictCheck = false, bool CheckType = true,
+                bool ExpensivePtrCheck = false);
 
 /// Attempt to sort the pointers in \p VL and return the sorted indices
 /// in \p SortedIndices, if reordering is required.
diff --git a/llvm/lib/Analysis/LoopAccessAnalysis.cpp b/llvm/lib/Analysis/LoopAccessAnalysis.cpp
index 409d9ceb5b812..d33bf13304c13 100644
--- a/llvm/lib/Analysis/LoopAccessAnalysis.cpp
+++ b/llvm/lib/Analysis/LoopAccessAnalysis.cpp
@@ -1806,7 +1806,8 @@ std::optional<int64_t> llvm::getPointersDiff(Type *ElemTyA, Value *PtrA,
                                              Type *ElemTyB, Value *PtrB,
                                              const DataLayout &DL,
                                              ScalarEvolution &SE,
-                                             bool StrictCheck, bool CheckType) {
+                                             bool StrictCheck, bool CheckType,
+                                             bool ExpensivePtrCheck) {
   assert(PtrA && PtrB && "Expected non-nullptr pointers.");
 
   // Make sure that A and B are different pointers.
@@ -1851,8 +1852,19 @@ std::optional<int64_t> llvm::getPointersDiff(Type *ElemTyA, Value *PtrA,
     // Otherwise compute the distance with SCEV between the base pointers.
     const SCEV *PtrSCEVA = SE.getSCEV(PtrA);
     const SCEV *PtrSCEVB = SE.getSCEV(PtrB);
-    std::optional<APInt> Diff =
-        SE.computeConstantDifference(PtrSCEVB, PtrSCEVA);
+    std::optional<APInt> Diff;
+    if (ExpensivePtrCheck) {
+      const SCEV *MinusSCEV = SE.getMinusSCEV(PtrSCEVB, PtrSCEVA);
+      if (MinusSCEV == SE.getCouldNotCompute())
+        return std::nullopt;
+      ConstantRange DistRange = SE.getSignedRange(MinusSCEV);
+      if (!DistRange.isSingleElement())
+        return std::nullopt;
+      APInt Dist = DistRange.getSingleElement()->sextOrTrunc(IdxWidth);
+      Diff = (OffsetB - OffsetA + Dist).sextOrTrunc(IdxWidth);
+    } else {
+      Diff = SE.computeConstantDifference(PtrSCEVB, PtrSCEVA);
+    }
     if (!Diff)
       return std::nullopt;
     Val = Diff->trySExtValue();
diff --git a/llvm/unittests/Analysis/CMakeLists.txt b/llvm/unittests/Analysis/CMakeLists.txt
index 9ca57024a6134..c07f3cc3a16bf 100644
--- a/llvm/unittests/Analysis/CMakeLists.txt
+++ b/llvm/unittests/Analysis/CMakeLists.txt
@@ -38,6 +38,7 @@ set(ANALYSIS_TEST_SOURCES
   LastRunTrackingAnalysisTest.cpp
   LazyCallGraphTest.cpp
   LoadsTest.cpp
+  LoopAccessAnalysisTest.cpp
   LoopInfoTest.cpp
   LoopNestTest.cpp
   MemoryBuiltinsTest.cpp
diff --git a/llvm/unittests/Analysis/LoopAccessAnalysisTest.cpp b/llvm/unittests/Analysis/LoopAccessAnalysisTest.cpp
new file mode 100644
index 0000000000000..b7411c8ec4bec
--- /dev/null
+++ b/llvm/unittests/Analysis/LoopAccessAnalysisTest.cpp
@@ -0,0 +1,84 @@
+//===- LoopAccessAnalysisTest.cpp - LoopAccessAnalysis unit 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/LoopAccessAnalysis.h"
+#include "llvm/Analysis/AssumptionCache.h"
+#include "llvm/Analysis/LoopInfo.h"
+#include "llvm/Analysis/ScalarEvolution.h"
+#include "llvm/Analysis/TargetLibraryInfo.h"
+#include "llvm/AsmParser/Parser.h"
+#include "llvm/IR/Dominators.h"
+#include "llvm/IR/LLVMContext.h"
+#include "llvm/IR/Module.h"
+#include "llvm/Support/SourceMgr.h"
+#include "gtest/gtest.h"
+
+using namespace llvm;
+
+class LoopAccessAnalysisTest : public testing::Test {
+protected:
+  LLVMContext Context;
+  std::unique_ptr<Module> M;
+
+  void parseIR(StringRef Assembly) {
+    SMDiagnostic Error;
+    M = parseAssemblyString(Assembly, Error, Context);
+    if (!M) {
+      std::string Msg;
+      raw_string_ostream OS(Msg);
+      Error.print("", OS);
+      report_fatal_error(Twine(Msg));
+    }
+  }
+};
+
+// A simple test to exercise the ExpensivePtrCheck argument. It should return
+// the same value for simple pointer expressions like in this test case.
+TEST_F(LoopAccessAnalysisTest, GetPointersDiff_ExpensivePtrCheck) {
+  parseIR(R"(
+    define void @foo(ptr %ptr0, ptr %ptr1) {
+      %gep0 = getelementptr i8, ptr %ptr0, i64 0
+      %gep1 = getelementptr i8, ptr %ptr0, i64 42
+      ret void
+    }
+  )");
+
+  Function &F = *M->getFunction("foo");
+  TargetLibraryInfoImpl TLII(M->getTargetTriple());
+  TargetLibraryInfo TLI(TLII);
+  AssumptionCache AC(F);
+  DominatorTree DT(F);
+  LoopInfo LI(DT);
+  ScalarEvolution SE(F, TLI, AC, DT, LI);
+  const DataLayout &DL = M->getDataLayout();
+
+  BasicBlock &BB = *F.begin();
+  auto It = BB.begin();
+  Instruction *PtrA = &*It++;
+  Instruction *PtrB = &*It++;
+  Type *ElemTy = IntegerType::getInt8Ty(M->getContext());
+  auto Diff_F = getPointersDiff(ElemTy, PtrA, ElemTy, PtrB, DL, SE,
+                                /*StrictCheck=*/false, /*CheckType=*/true,
+                                /*ExpensivePtrCheck=*/false);
+  auto Diff_T = getPointersDiff(ElemTy, PtrA, ElemTy, PtrB, DL, SE,
+                                /*StrictCheck=*/false, /*CheckType=*/true,
+                                /*ExpensivePtrCheck=*/true);
+  EXPECT_EQ(Diff_F, Diff_T);
+  EXPECT_EQ(Diff_F, 42);
+
+  auto *Ptr0 = F.getArg(0);
+  auto *Ptr1 = F.getArg(1);
+  Diff_F = getPointersDiff(ElemTy, Ptr0, ElemTy, Ptr1, DL, SE,
+                           /*StrictCheck=*/false, /*CheckType=*/true,
+                           /*ExpensivePtrCheck=*/false);
+  Diff_T = getPointersDiff(ElemTy, Ptr0, ElemTy, Ptr1, DL, SE,
+                           /*StrictCheck=*/false, /*CheckType=*/true,
+                           /*ExpensivePtrCheck=*/true);
+  EXPECT_EQ(Diff_F, Diff_T);
+  EXPECT_EQ(Diff_F, std::nullopt);
+}

``````````

</details>


https://github.com/llvm/llvm-project/pull/226333


More information about the llvm-commits mailing list