[llvm] [SimplifyQuery] Include a context-function (NFC) (PR #225369)

Ramkumar Ramachandra via llvm-commits llvm-commits at lists.llvm.org
Tue Sep 22 04:19:47 PDT 2026


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

The context-function logically belongs in SimplifyQuery, and this helps simplify some code, and allows us to represent a context-function in the absence of a context-instruction.

>From 239a5fdbe1910eaacd8c89923ebd8ae8256f1609 Mon Sep 17 00:00:00 2001
From: Ramkumar Ramachandra <artagnon at tenstorrent.com>
Date: Tue, 22 Sep 2026 11:21:20 +0100
Subject: [PATCH] [SimplifyQuery] Include a context-function (NFC)

The context-function logically belongs in SimplifyQuery, and this helps
simplify some code, and allows us to represent a context-function in the
absence of a context-instruction.
---
 llvm/include/llvm/Analysis/ConstantFolding.h  |  6 +--
 .../llvm/Analysis/InstSimplifyFolder.h        |  2 +-
 .../llvm/Analysis/InstructionSimplify.h       |  1 -
 llvm/include/llvm/Analysis/SimplifyQuery.h    | 23 +++++++--
 llvm/lib/Analysis/ConstantFolding.cpp         | 49 ++++++++++---------
 llvm/lib/Analysis/InstructionSimplify.cpp     | 40 +++++++--------
 6 files changed, 67 insertions(+), 54 deletions(-)

diff --git a/llvm/include/llvm/Analysis/ConstantFolding.h b/llvm/include/llvm/Analysis/ConstantFolding.h
index ddbeecf0b43f0..72a11e5c94469 100644
--- a/llvm/include/llvm/Analysis/ConstantFolding.h
+++ b/llvm/include/llvm/Analysis/ConstantFolding.h
@@ -87,7 +87,7 @@ LLVM_ABI Constant *ConstantFoldInstOperands(
 /// Denormal inputs may be flushed based on the denormal handling mode.
 LLVM_ABI Constant *ConstantFoldCompareInstOperands(
     unsigned Predicate, Constant *LHS, Constant *RHS, const DataLayout &DL,
-    const TargetLibraryInfo *TLI = nullptr, const Instruction *I = nullptr);
+    const TargetLibraryInfo *TLI = nullptr, const Function *CxtF = nullptr);
 
 /// Attempt to constant fold a unary operation with the specified operand.
 /// Returns null on failure.
@@ -116,7 +116,7 @@ ConstantFoldFPInstOperands(unsigned Opcode, Constant *LHS, Constant *RHS,
 ///
 /// If the calling function's denormal_fpenv input mode is dynamic for the
 /// floating-point type, returns nullptr for denormal inputs.
-LLVM_ABI Constant *FlushFPConstant(Constant *Operand, const Instruction *I,
+LLVM_ABI Constant *FlushFPConstant(Constant *Operand, const Function *CxtF,
                                    bool IsOutput);
 
 /// Attempt to constant fold a cast with the specified operand.  If it
@@ -173,7 +173,7 @@ LLVM_ABI Constant *ConstantFoldCall(const CallBase *Call, Function *F,
 LLVM_ABI Constant *ConstantFoldIntrinsic(Intrinsic::ID ID,
                                          ArrayRef<Constant *> Ops, Type *Ty,
                                          const DataLayout &DL,
-                                         Function *CxtF = nullptr);
+                                         const Function *CxtF = nullptr);
 
 /// ConstantFoldLoadThroughBitcast - try to cast constant to destination type
 /// returning null if unsuccessful. Can cast pointer to pointer or pointer to
diff --git a/llvm/include/llvm/Analysis/InstSimplifyFolder.h b/llvm/include/llvm/Analysis/InstSimplifyFolder.h
index 16286b93677f8..6414901a4f6ed 100644
--- a/llvm/include/llvm/Analysis/InstSimplifyFolder.h
+++ b/llvm/include/llvm/Analysis/InstSimplifyFolder.h
@@ -123,7 +123,7 @@ class LLVM_ABI InstSimplifyFolder final : public IRBuilderFolder {
   Value *FoldIntrinsic(Intrinsic::ID ID, ArrayRef<Value *> Ops, Type *Ty,
                        FastMathFlags FMF = {},
                        Function *CtxF = nullptr) const override {
-    return simplifyIntrinsic(ID, Ty, Ops, FMF, SQ, CtxF);
+    return simplifyIntrinsic(ID, Ty, Ops, FMF, SQ.getWithFunction(CtxF));
   }
 
   //===--------------------------------------------------------------------===//
diff --git a/llvm/include/llvm/Analysis/InstructionSimplify.h b/llvm/include/llvm/Analysis/InstructionSimplify.h
index 61e1c59b59f56..db73c2d78eeda 100644
--- a/llvm/include/llvm/Analysis/InstructionSimplify.h
+++ b/llvm/include/llvm/Analysis/InstructionSimplify.h
@@ -204,7 +204,6 @@ LLVM_ABI Value *simplifyAddrSpaceCastInst(Value *Op, Type *Ty, bool IsNonNull,
 LLVM_ABI Value *
 simplifyIntrinsic(Intrinsic::ID IID, Type *ReturnType, ArrayRef<Value *> Args,
                   FastMathFlags FMF, const SimplifyQuery &Q,
-                  Function *CxtF = nullptr,
                   fp::ExceptionBehavior ExBehavior = fp::ebIgnore,
                   RoundingMode Rounding = RoundingMode::NearestTiesToEven);
 
diff --git a/llvm/include/llvm/Analysis/SimplifyQuery.h b/llvm/include/llvm/Analysis/SimplifyQuery.h
index b81b1dae27471..88c399205c835 100644
--- a/llvm/include/llvm/Analysis/SimplifyQuery.h
+++ b/llvm/include/llvm/Analysis/SimplifyQuery.h
@@ -76,6 +76,7 @@ struct SimplifyQuery {
   const Instruction *CxtI = nullptr;
   const DomConditionCache *DC = nullptr;
   const CondContext *CC = nullptr;
+  const Function *CxtF = nullptr;
 
   // Wrapper to query additional information for instructions like metadata or
   // keywords like nsw, which provides conservative results if those cannot
@@ -89,7 +90,10 @@ struct SimplifyQuery {
   bool AllowEphemerals = false;
 
   SimplifyQuery(const DataLayout &DL, const Instruction *CXTI = nullptr)
-      : DL(DL), CxtI(CXTI) {}
+      : DL(DL), CxtI(CXTI) {
+    if (CxtI)
+      CxtF = CxtI->getFunction();
+  }
 
   SimplifyQuery(const DataLayout &DL, const TargetLibraryInfo *TLI,
                 const DominatorTree *DT = nullptr,
@@ -97,18 +101,31 @@ struct SimplifyQuery {
                 const Instruction *CXTI = nullptr, bool UseInstrInfo = true,
                 bool CanUseUndef = true, const DomConditionCache *DC = nullptr)
       : DL(DL), TLI(TLI), DT(DT), AC(AC), CxtI(CXTI), DC(DC), IIQ(UseInstrInfo),
-        CanUseUndef(CanUseUndef) {}
+        CanUseUndef(CanUseUndef) {
+    if (CxtI)
+      CxtF = CxtI->getFunction();
+  }
 
   SimplifyQuery(const DataLayout &DL, const DominatorTree *DT,
                 AssumptionCache *AC = nullptr,
                 const Instruction *CXTI = nullptr, bool UseInstrInfo = true,
                 bool CanUseUndef = true)
       : DL(DL), DT(DT), AC(AC), CxtI(CXTI), IIQ(UseInstrInfo),
-        CanUseUndef(CanUseUndef) {}
+        CanUseUndef(CanUseUndef) {
+    if (CxtI)
+      CxtF = CxtI->getFunction();
+  }
 
   SimplifyQuery getWithInstruction(const Instruction *I) const {
     SimplifyQuery Copy(*this);
     Copy.CxtI = I;
+    if (I)
+      Copy.CxtF = I->getFunction();
+    return Copy;
+  }
+  SimplifyQuery getWithFunction(const Function *F) const {
+    SimplifyQuery Copy(*this);
+    Copy.CxtF = F;
     return Copy;
   }
   SimplifyQuery getWithoutUndef() const {
diff --git a/llvm/lib/Analysis/ConstantFolding.cpp b/llvm/lib/Analysis/ConstantFolding.cpp
index b83c9bb13ecb4..a9dedc622b998 100644
--- a/llvm/lib/Analysis/ConstantFolding.cpp
+++ b/llvm/lib/Analysis/ConstantFolding.cpp
@@ -1181,7 +1181,7 @@ Constant *ConstantFoldInstOperandsImpl(const Value *InstOrCE, unsigned Opcode,
   case Instruction::FCmp: {
     auto *C = cast<CmpInst>(InstOrCE);
     return ConstantFoldCompareInstOperands(C->getPredicate(), Ops[0], Ops[1],
-                                           DL, TLI, C);
+                                           DL, TLI, C->getFunction());
   }
   case Instruction::Freeze:
     return isGuaranteedNotToBeUndefOrPoison(Ops[0]) ? Ops[0] : nullptr;
@@ -1327,9 +1327,11 @@ Constant *llvm::ConstantFoldInstOperands(const Instruction *I,
                                       AllowNonDeterministic);
 }
 
-Constant *llvm::ConstantFoldCompareInstOperands(
-    unsigned IntPredicate, Constant *Ops0, Constant *Ops1, const DataLayout &DL,
-    const TargetLibraryInfo *TLI, const Instruction *I) {
+Constant *llvm::ConstantFoldCompareInstOperands(unsigned IntPredicate,
+                                                Constant *Ops0, Constant *Ops1,
+                                                const DataLayout &DL,
+                                                const TargetLibraryInfo *TLI,
+                                                const Function *CxtF) {
   CmpInst::Predicate Predicate = (CmpInst::Predicate)IntPredicate;
   // fold: icmp (inttoptr x), null         -> icmp x, 0
   // fold: icmp null, (inttoptr x)         -> icmp 0, x
@@ -1430,10 +1432,10 @@ Constant *llvm::ConstantFoldCompareInstOperands(
   if (CmpInst::isFPPredicate(Predicate)) {
     // Flush any denormal constant float input according to denormal handling
     // mode.
-    Ops0 = FlushFPConstant(Ops0, I, /*IsOutput=*/false);
+    Ops0 = FlushFPConstant(Ops0, CxtF, /*IsOutput=*/false);
     if (!Ops0)
       return nullptr;
-    Ops1 = FlushFPConstant(Ops1, I, /*IsOutput=*/false);
+    Ops1 = FlushFPConstant(Ops1, CxtF, /*IsOutput=*/false);
     if (!Ops1)
       return nullptr;
   }
@@ -1482,29 +1484,27 @@ static ConstantFP *flushDenormalConstant(Type *Ty, const APFloat &APF,
 
 /// Return the denormal mode that can be assumed when executing a floating point
 /// operation at \p CtxI.
-static DenormalMode getInstrDenormalMode(const Instruction *CtxI, Type *Ty) {
-  if (!CtxI || !CtxI->getParent() || !CtxI->getFunction())
+static DenormalMode getInstrDenormalMode(const Function *CtxF, Type *Ty) {
+  if (!CtxF)
     return DenormalMode::getDynamic();
-  return CtxI->getFunction()->getDenormalMode(
-      Ty->getScalarType()->getFltSemantics());
+  return CtxF->getDenormalMode(Ty->getScalarType()->getFltSemantics());
 }
 
-static ConstantFP *flushDenormalConstantFP(ConstantFP *CFP,
-                                           const Instruction *Inst,
-                                           bool IsOutput) {
+static ConstantFP *
+flushDenormalConstantFP(ConstantFP *CFP, const Function *CxtF, bool IsOutput) {
   const APFloat &APF = CFP->getValueAPF();
   if (!APF.isDenormal())
     return CFP;
 
-  DenormalMode Mode = getInstrDenormalMode(Inst, CFP->getType());
+  DenormalMode Mode = getInstrDenormalMode(CxtF, CFP->getType());
   return flushDenormalConstant(CFP->getType(), APF,
                                IsOutput ? Mode.Output : Mode.Input);
 }
 
-Constant *llvm::FlushFPConstant(Constant *Operand, const Instruction *Inst,
+Constant *llvm::FlushFPConstant(Constant *Operand, const Function *CxtF,
                                 bool IsOutput) {
   if (ConstantFP *CFP = dyn_cast<ConstantFP>(Operand))
-    return flushDenormalConstantFP(CFP, Inst, IsOutput);
+    return flushDenormalConstantFP(CFP, CxtF, IsOutput);
 
   if (isa<ConstantAggregateZero, UndefValue>(Operand))
     return Operand;
@@ -1513,7 +1513,7 @@ Constant *llvm::FlushFPConstant(Constant *Operand, const Instruction *Inst,
   VectorType *VecTy = dyn_cast<VectorType>(Ty);
   if (VecTy) {
     if (auto *Splat = dyn_cast_or_null<ConstantFP>(Operand->getSplatValue())) {
-      ConstantFP *Folded = flushDenormalConstantFP(Splat, Inst, IsOutput);
+      ConstantFP *Folded = flushDenormalConstantFP(Splat, CxtF, IsOutput);
       if (!Folded)
         return nullptr;
       return ConstantVector::getSplat(VecTy->getElementCount(), Folded);
@@ -1538,7 +1538,7 @@ Constant *llvm::FlushFPConstant(Constant *Operand, const Instruction *Inst,
       if (!CFP)
         return nullptr;
 
-      ConstantFP *Folded = flushDenormalConstantFP(CFP, Inst, IsOutput);
+      ConstantFP *Folded = flushDenormalConstantFP(CFP, CxtF, IsOutput);
       if (!Folded)
         return nullptr;
       NewElts.push_back(Folded);
@@ -1554,7 +1554,7 @@ Constant *llvm::FlushFPConstant(Constant *Operand, const Instruction *Inst,
       if (!Elt.isDenormal()) {
         NewElts.push_back(ConstantFP::get(Ty, Elt));
       } else {
-        DenormalMode Mode = getInstrDenormalMode(Inst, Ty);
+        DenormalMode Mode = getInstrDenormalMode(CxtF, Ty);
         ConstantFP *Folded =
             flushDenormalConstant(Ty, Elt, IsOutput ? Mode.Output : Mode.Input);
         if (!Folded)
@@ -1575,10 +1575,12 @@ Constant *llvm::ConstantFoldFPInstOperands(unsigned Opcode, Constant *LHS,
                                            bool AllowNonDeterministic) {
   if (Instruction::isBinaryOp(Opcode)) {
     // Flush denormal inputs if needed.
-    Constant *Op0 = FlushFPConstant(LHS, I, /* IsOutput */ false);
+    Constant *Op0 =
+        FlushFPConstant(LHS, I->getFunction(), /* IsOutput */ false);
     if (!Op0)
       return nullptr;
-    Constant *Op1 = FlushFPConstant(RHS, I, /* IsOutput */ false);
+    Constant *Op1 =
+        FlushFPConstant(RHS, I->getFunction(), /* IsOutput */ false);
     if (!Op1)
       return nullptr;
 
@@ -1597,7 +1599,7 @@ Constant *llvm::ConstantFoldFPInstOperands(unsigned Opcode, Constant *LHS,
       return nullptr;
 
     // Flush denormal output if needed.
-    C = FlushFPConstant(C, I, /* IsOutput */ true);
+    C = FlushFPConstant(C, I->getFunction(), /* IsOutput */ true);
     if (!C)
       return nullptr;
 
@@ -4786,7 +4788,8 @@ ConstantFoldStructCall(StringRef Name, Intrinsic::ID IntrinsicID,
 
 Constant *llvm::ConstantFoldIntrinsic(Intrinsic::ID ID,
                                       ArrayRef<Constant *> Ops, Type *Ty,
-                                      const DataLayout &DL, Function *CxtF) {
+                                      const DataLayout &DL,
+                                      const Function *CxtF) {
   // In the absence of CxtF, assume strictfp conservatively.
   if (!canConstantFoldIntrinsic(ID, CxtF ? CxtF->isStrictFP() : true) ||
       (DisableFPCallFolding &&
diff --git a/llvm/lib/Analysis/InstructionSimplify.cpp b/llvm/lib/Analysis/InstructionSimplify.cpp
index d9aebc757d636..9a810c93842de 100644
--- a/llvm/lib/Analysis/InstructionSimplify.cpp
+++ b/llvm/lib/Analysis/InstructionSimplify.cpp
@@ -4204,7 +4204,7 @@ static Value *simplifyFCmpInst(CmpPredicate Pred, Value *LHS, Value *RHS,
     if (Constant *CRHS = dyn_cast<Constant>(RHS)) {
       // if the folding isn't successfull, fall back to the rest of the logic
       if (auto *Result = ConstantFoldCompareInstOperands(Pred, CLHS, CRHS, Q.DL,
-                                                         Q.TLI, Q.CxtI))
+                                                         Q.TLI, Q.CxtF))
         return Result;
     } else {
       // If we have a constant, make sure it is on the RHS.
@@ -4274,14 +4274,9 @@ static Value *simplifyFCmpInst(CmpPredicate Pred, Value *LHS, Value *RHS,
     return computeKnownFPClass(LHS, FMF, InterestedFlags, Q);
   };
 
-  if (C && Q.CxtI) {
+  if (C && Q.CxtF) {
     // Fold out compares that express a class test.
-    //
-    // FIXME: Should be able to perform folds without context
-    // instruction. Always pass in the context function?
-
-    const Function *ParentF = Q.CxtI->getFunction();
-    auto [ClassVal, ClassTest] = fcmpToClassTest(Pred, *ParentF, LHS, C);
+    auto [ClassVal, ClassTest] = fcmpToClassTest(Pred, *Q.CxtF, LHS, C);
     if (ClassVal) {
       FullKnownClassLHS = computeLHSClass();
       if ((FullKnownClassLHS->getKnownFPClasses() & ClassTest) == fcNone)
@@ -5730,9 +5725,8 @@ Value *llvm::simplifyCastInst(unsigned CastOpc, Value *Op, Type *Ty,
 static Value *simplifyAddrSpaceCastInst(Value *Op, Type *Ty, bool IsNonNull,
                                         const SimplifyQuery &Q,
                                         unsigned MaxRecurse) {
-  if (IsNonNull && isa<ConstantPointerNull>(Op) && Q.CxtI &&
-      !NullPointerIsDefined(Q.CxtI->getFunction(),
-                            Op->getType()->getPointerAddressSpace()))
+  if (IsNonNull && isa<ConstantPointerNull>(Op) && Q.CxtF &&
+      !NullPointerIsDefined(Q.CxtF, Op->getType()->getPointerAddressSpace()))
     return PoisonValue::get(Ty);
 
   return ::simplifyCastInst(Instruction::AddrSpaceCast, Op, Ty, Q, MaxRecurse);
@@ -6956,10 +6950,10 @@ static Value *simplifyBinaryIntrinsic(Intrinsic::ID IID, Type *ReturnType,
     if (match(Op1, m_Zero()))
       return ConstantInt::getFalse(ReturnType);
 
-    if (!Q.CxtI)
+    if (!Q.CxtF)
       break;
 
-    const Function *F = Q.CxtI->getFunction();
+    const Function *F = Q.CxtF;
     auto *ScalableTy = dyn_cast<ScalableVectorType>(ReturnType);
     Attribute Attr = F->getFnAttribute(Attribute::VScaleRange);
     if (ScalableTy && Attr.isValid()) {
@@ -7376,7 +7370,7 @@ static Value *simplifyIdentityInterleave(Intrinsic::ID IID,
 
 Value *llvm::simplifyIntrinsic(Intrinsic::ID IID, Type *ReturnType,
                                ArrayRef<Value *> Args, FastMathFlags FMF,
-                               const SimplifyQuery &Q, Function *CxtF,
+                               const SimplifyQuery &Q,
                                fp::ExceptionBehavior ExBehavior,
                                RoundingMode Rounding) {
   unsigned NumOperands = Args.size();
@@ -7388,7 +7382,7 @@ Value *llvm::simplifyIntrinsic(Intrinsic::ID IID, Type *ReturnType,
   if (all_of(Args, IsaPred<Constant>))
     if (Constant *C = ConstantFoldIntrinsic(
             IID, ArrayRef((Constant *const *)Args.data(), Args.size()),
-            ReturnType, Q.DL, CxtF))
+            ReturnType, Q.DL, Q.CxtF))
       return C;
 
   // Most of the intrinsics with no operands have some kind of side effect.
@@ -7396,9 +7390,9 @@ Value *llvm::simplifyIntrinsic(Intrinsic::ID IID, Type *ReturnType,
   if (!NumOperands) {
     switch (IID) {
     case Intrinsic::vscale: {
-      if (!CxtF)
+      if (!Q.CxtF)
         return nullptr;
-      ConstantRange CR = getVScaleRange(CxtF, 64);
+      ConstantRange CR = getVScaleRange(Q.CxtF, 64);
       if (const APInt *C = CR.getSingleElement())
         return ConstantInt::get(ReturnType, C->getZExtValue());
       return nullptr;
@@ -7540,8 +7534,8 @@ Value *llvm::simplifyIntrinsic(Intrinsic::ID IID, Type *ReturnType,
     ConstantRange NumElts(
         APInt(BitWidth, Ty->getElementCount().getKnownMinValue()));
     if (Ty->isScalableTy())
-      NumElts = NumElts.multiply(CxtF ? getVScaleRange(CxtF, BitWidth)
-                                      : ConstantRange::getFull(BitWidth));
+      NumElts = NumElts.multiply(Q.CxtF ? getVScaleRange(Q.CxtF, BitWidth)
+                                        : ConstantRange::getFull(BitWidth));
 
     // If we know Offset > NumElts, simplify to poison.
     ConstantRange CR = computeConstantRangeIncludingKnownBits(Offset, false, Q);
@@ -7623,9 +7617,9 @@ static Value *simplifyIntrinsic(CallBase *Call, ArrayRef<Value *> Args,
       ExBehavior = Constrained->getExceptionBehavior().value_or(ExBehavior);
       Rounding = Constrained->getRoundingMode().value_or(Rounding);
     }
-    return simplifyIntrinsic(IID, ReturnType, Args,
-                             Call->getFastMathFlagsOrNone(), Q,
-                             Call->getFunction(), ExBehavior, Rounding);
+    return simplifyIntrinsic(
+        IID, ReturnType, Args, Call->getFastMathFlagsOrNone(),
+        Q.getWithFunction(Call->getFunction()), ExBehavior, Rounding);
   }
   }
 }
@@ -7744,7 +7738,7 @@ static Value *simplifyInstructionWithOperands(Instruction *I,
                                               const SimplifyQuery &SQ,
                                               unsigned MaxRecurse) {
   assert(I->getFunction() && "instruction should be inserted in a function");
-  assert((!SQ.CxtI || SQ.CxtI->getFunction() == I->getFunction()) &&
+  assert((!SQ.CxtF || SQ.CxtF == I->getFunction()) &&
          "context instruction should be in the same function");
 
   const SimplifyQuery Q = SQ.CxtI ? SQ : SQ.getWithInstruction(I);



More information about the llvm-commits mailing list