[llvm] Add MUL/SHL handling in decomposeLinearExpression (PR #222574)
Arne Stenkrona via llvm-commits
llvm-commits at lists.llvm.org
Thu Sep 10 03:02:10 PDT 2026
https://github.com/ArneStenkrona2 created https://github.com/llvm/llvm-project/pull/222574
Implements a FIXME in decomposeLinearExpression, allowing it to look through multiplications and left-shifts.
We peel the variable index and keep track of the cumulated scale from the multiplications and shifts and apply this scale when we return the final expression.
This allows decomposeLinearExpression to extract more statically known information into the `Scale` of the linear expression.
assisted-by: codex
>From 25d38bb10fd22b6d05ec1a27eb68e1c3a7ac03d7 Mon Sep 17 00:00:00 2001
From: Arne Stenkrona <arne.stenkrona at arm.com>
Date: Mon, 7 Sep 2026 17:03:41 +0200
Subject: [PATCH] Add MUL/SHL handling in decomposeLinearExpression
Implements a FIXME in decomposeLinearExpression, allowing it to look
through multiplications and left-shifts.
We peel the variable index and keep track of the cumulated scale from
the multiplications and shifts and apply this scale when we return the
final expression.
This allows decomposeLinearExpression to extract more statically known
information into the `Scale` of the linear expression.
assisted-by: codex
---
llvm/lib/Analysis/Loads.cpp | 28 +++++++++++++++--
llvm/unittests/Analysis/LoadsTest.cpp | 44 +++++++++++++++++++++++++++
2 files changed, 69 insertions(+), 3 deletions(-)
diff --git a/llvm/lib/Analysis/Loads.cpp b/llvm/lib/Analysis/Loads.cpp
index de9022c540d42..cd5aede82385b 100644
--- a/llvm/lib/Analysis/Loads.cpp
+++ b/llvm/lib/Analysis/Loads.cpp
@@ -24,8 +24,10 @@
#include "llvm/IR/GetElementPtrTypeIterator.h"
#include "llvm/IR/IntrinsicInst.h"
#include "llvm/IR/Operator.h"
+#include "llvm/IR/PatternMatch.h"
using namespace llvm;
+using namespace llvm::PatternMatch;
static bool isAligned(const Value *Base, Align Alignment,
const DataLayout &DL) {
@@ -935,6 +937,26 @@ LinearExpression llvm::decomposeLinearExpression(const DataLayout &DL,
VarIndex = Index;
}
+ APInt IndexScale(BitWidth, 1);
+ while (auto *BO = dyn_cast_or_null<BinaryOperator>(VarIndex)) {
+ Value *UnscaledIndex = nullptr;
+ APInt Factor(BitWidth, 0);
+ ConstantInt *Constant;
+ if (match(BO, m_c_Mul(m_Value(UnscaledIndex), m_ConstantInt(Constant)))) {
+ Factor = Constant->getValue().zextOrTrunc(BitWidth);
+ } else if (match(BO, m_Shl(m_Value(UnscaledIndex),
+ m_ConstantInt(Constant)))) {
+ if (Constant->getValue().uge(BitWidth))
+ break;
+ Factor.setBit(Constant->getZExtValue());
+ } else {
+ break;
+ }
+
+ IndexScale *= Factor;
+ VarIndex = UnscaledIndex;
+ }
+
// Don't return non-canonical indexes.
if (VarIndex && !VarIndex->getType()->isIntegerTy(BitWidth))
return Expr;
@@ -963,12 +985,12 @@ LinearExpression llvm::decomposeLinearExpression(const DataLayout &DL,
continue;
}
- // FIXME: Also look through a mul/shl in the index.
assert(Expr.Index == nullptr && "Shouldn't have index yet");
- Expr.Index = Index;
+ Expr.Index = VarIndex;
// Truncate if type size exceeds index space.
Expr.Scale = APInt(BitWidth, GTI.getSequentialElementStride(DL),
- /*isSigned=*/false, /*implicitTrunc=*/true);
+ /*isSigned=*/false, /*implicitTrunc=*/true) *
+ IndexScale;
}
}
diff --git a/llvm/unittests/Analysis/LoadsTest.cpp b/llvm/unittests/Analysis/LoadsTest.cpp
index 8b15bda08485e..a5aa27715e4a2 100644
--- a/llvm/unittests/Analysis/LoadsTest.cpp
+++ b/llvm/unittests/Analysis/LoadsTest.cpp
@@ -17,6 +17,7 @@
#include "llvm/IR/Instructions.h"
#include "llvm/IR/LLVMContext.h"
#include "llvm/IR/Module.h"
+#include "llvm/IR/ValueSymbolTable.h"
#include "llvm/Support/SourceMgr.h"
#include "gtest/gtest.h"
@@ -243,3 +244,46 @@ loop.end:
ASSERT_TRUE((NonDerefLoads.size() == 1) &&
(NonDerefLoads[0]->getName() == "ld1"));
}
+
+TEST(LoadsTest, DecomposeLinearExpressionScaling) {
+ LLVMContext C;
+ std::unique_ptr<Module> M = parseIR(C, R"IR(
+target datalayout = "e-p:64:64"
+
+define void @f(ptr %base, i64 %index) {
+ %mul = mul i64 %index, 3
+ %shl = shl i64 %index, 2
+ %scaled = shl i64 %mul, 2
+ %gep.mul = getelementptr i32, ptr %base, i64 %mul
+ %gep.shl = getelementptr i32, ptr %base, i64 %shl
+ %gep.scaled = getelementptr i32, ptr %base, i64 %scaled
+ ret void
+}
+)IR");
+ ASSERT_TRUE(M);
+
+ Function *F = M->getFunction("f");
+ ASSERT_TRUE(F);
+ Value *Base = F->getArg(0);
+ Value *Index = F->getArg(1);
+ Value *GEPMul = F->getValueSymbolTable()->lookup("gep.mul");
+ Value *GEPShl = F->getValueSymbolTable()->lookup("gep.shl");
+ Value *GEPScaled = F->getValueSymbolTable()->lookup("gep.scaled");
+ ASSERT_TRUE(GEPMul && GEPShl && GEPScaled);
+
+ const DataLayout &DL = M->getDataLayout();
+ LinearExpression Expr = decomposeLinearExpression(DL, GEPMul);
+ EXPECT_EQ(Base, Expr.BasePtr);
+ EXPECT_EQ(Index, Expr.Index);
+ EXPECT_EQ(APInt(64, 12), Expr.Scale);
+
+ Expr = decomposeLinearExpression(DL, GEPShl);
+ EXPECT_EQ(Base, Expr.BasePtr);
+ EXPECT_EQ(Index, Expr.Index);
+ EXPECT_EQ(APInt(64, 16), Expr.Scale);
+
+ Expr = decomposeLinearExpression(DL, GEPScaled);
+ EXPECT_EQ(Base, Expr.BasePtr);
+ EXPECT_EQ(Index, Expr.Index);
+ EXPECT_EQ(APInt(64, 48), Expr.Scale);
+}
More information about the llvm-commits
mailing list