[llvm] 05f6994 - [IR][NFC] Introduce Constant::containsMatchingVectorElement and corresponding matcher (#200502)

via llvm-commits llvm-commits at lists.llvm.org
Sat Jun 27 03:40:07 PDT 2026


Author: Sean Clarke
Date: 2026-06-27T11:40:03+01:00
New Revision: 05f69942c5025295a4ddcab2b06d50774ac8f625

URL: https://github.com/llvm/llvm-project/commit/05f69942c5025295a4ddcab2b06d50774ac8f625
DIFF: https://github.com/llvm/llvm-project/commit/05f69942c5025295a4ddcab2b06d50774ac8f625.diff

LOG: [IR][NFC] Introduce Constant::containsMatchingVectorElement and corresponding matcher (#200502)

A common pattern when dealing with vectors is to iterate over each
element and check whether any element satisfies a condition. Introduce
`Constant::containsMatchingVectorElement` to generalize this behavior
along with a corresponding matcher `m_ContainsMatchingVectorElement`
which checks whether any elements match the given subpattern.

Remove function `llvm::maskContainsAllOneOrUndef` in favor of using this
generalization instead.

Co-authored-by: Ramkumar Ramachandra <r at artagnon.com>

Added: 
    

Modified: 
    llvm/include/llvm/IR/Constant.h
    llvm/include/llvm/IR/PatternMatch.h
    llvm/lib/Analysis/VectorUtils.cpp
    llvm/lib/IR/Constants.cpp
    llvm/lib/Transforms/InstCombine/InstCombineMulDivRem.cpp
    llvm/lib/Transforms/InstCombine/InstructionCombining.cpp

Removed: 
    


################################################################################
diff  --git a/llvm/include/llvm/IR/Constant.h b/llvm/include/llvm/IR/Constant.h
index 2ab1400264327..1013b8dace9a2 100644
--- a/llvm/include/llvm/IR/Constant.h
+++ b/llvm/include/llvm/IR/Constant.h
@@ -130,6 +130,11 @@ class Constant : public User {
   /// any constant expressions.
   LLVM_ABI bool containsConstantExpression() const;
 
+  /// Return true if this is a vector constant where at least one element
+  /// satisfies the given predicate. Scalable vectors are not checked.
+  LLVM_ABI bool
+  containsMatchingVectorElement(function_ref<bool(Constant *)> PredFn) const;
+
   /// Return true if the value can vary between threads.
   LLVM_ABI bool isThreadDependent() const;
 

diff  --git a/llvm/include/llvm/IR/PatternMatch.h b/llvm/include/llvm/IR/PatternMatch.h
index 1b19175156261..95f5a9c5bba80 100644
--- a/llvm/include/llvm/IR/PatternMatch.h
+++ b/llvm/include/llvm/IR/PatternMatch.h
@@ -181,16 +181,32 @@ inline auto m_ConstantInt() { return m_Isa<ConstantInt>(); }
 /// Match an arbitrary ConstantFP and ignore it.
 inline auto m_ConstantFP() { return m_Isa<ConstantFP>(); }
 
-struct constantexpr_match {
+template <typename SPTy> struct ContainsMatchingVectorElement_match {
+  SPTy SubPattern;
+  ContainsMatchingVectorElement_match(const SPTy &SP) : SubPattern(SP) {}
+
   template <typename ITy> bool match(ITy *V) const {
     auto *C = dyn_cast<Constant>(V);
-    return C && (isa<ConstantExpr>(C) || C->containsConstantExpression());
+    return C && C->containsMatchingVectorElement(
+                    [&](Constant *E) { return SubPattern.match(E); });
   }
 };
 
+/// Match a vector constant where at least one of its elements matches the
+/// subpattern. Scalable vector constants are not matched. Any bindings in the
+/// subpattern will be bound to the first match.
+template <typename SPTy>
+inline ContainsMatchingVectorElement_match<SPTy>
+m_ContainsMatchingVectorElement(const SPTy &SubPattern) {
+  return SubPattern;
+}
+
 /// Match a constant expression or a constant that contains a constant
 /// expression.
-inline constantexpr_match m_ConstantExpr() { return constantexpr_match(); }
+inline auto m_ConstantExpr() {
+  return m_CombineOr(m_Isa<ConstantExpr>(),
+                     m_ContainsMatchingVectorElement(m_Isa<ConstantExpr>()));
+}
 
 template <typename SubPattern_t> struct Splat_match {
   SubPattern_t SubPattern;
@@ -888,13 +904,12 @@ inline match_bind<const BasicBlock> m_BasicBlock(const BasicBlock *&V) {
 struct immconstant_ty {
   template <typename ITy> static bool isImmConstant(ITy *V) {
     if (auto *CV = dyn_cast<Constant>(V)) {
-      if (!isa<ConstantExpr>(CV) && !CV->containsConstantExpression())
+      if (!match(CV, m_ConstantExpr()))
         return true;
 
       if (CV->getType()->isVectorTy()) {
         if (auto *Splat = CV->getSplatValue(/*AllowPoison=*/true)) {
-          if (!isa<ConstantExpr>(Splat) &&
-              !Splat->containsConstantExpression()) {
+          if (!match(Splat, m_ConstantExpr())) {
             return true;
           }
         }

diff  --git a/llvm/lib/Analysis/VectorUtils.cpp b/llvm/lib/Analysis/VectorUtils.cpp
index bcdb6df47d757..193fb6720cf60 100644
--- a/llvm/lib/Analysis/VectorUtils.cpp
+++ b/llvm/lib/Analysis/VectorUtils.cpp
@@ -1265,22 +1265,9 @@ bool llvm::maskContainsAllOneOrUndef(Value *Mask) {
              1 &&
          "Mask must be a vector of i1");
 
-  auto *ConstMask = dyn_cast<Constant>(Mask);
-  if (!ConstMask)
-    return false;
-  if (ConstMask->isAllOnesValue() || isa<UndefValue>(ConstMask))
-    return true;
-  if (isa<ScalableVectorType>(ConstMask->getType()))
-    return false;
-  for (unsigned
-           I = 0,
-           E = cast<FixedVectorType>(ConstMask->getType())->getNumElements();
-       I != E; ++I) {
-    if (auto *MaskElt = ConstMask->getAggregateElement(I))
-      if (MaskElt->isAllOnesValue() || isa<UndefValue>(MaskElt))
-        return true;
-  }
-  return false;
+  auto AllOneOrUndef = m_CombineOr(m_AllOnes(), m_UndefValue());
+  return match(Mask, m_CombineOr(AllOneOrUndef, m_ContainsMatchingVectorElement(
+                                                    AllOneOrUndef)));
 }
 
 /// TODO: This is a lot like known bits, but for

diff  --git a/llvm/lib/IR/Constants.cpp b/llvm/lib/IR/Constants.cpp
index 4327ad84874e6..a606a2a27669b 100644
--- a/llvm/lib/IR/Constants.cpp
+++ b/llvm/lib/IR/Constants.cpp
@@ -312,20 +312,13 @@ bool Constant::isElementWiseEqual(Value *Y) const {
 static bool
 containsUndefinedElement(const Constant *C,
                          function_ref<bool(const Constant *)> HasFn) {
-  if (auto *VTy = dyn_cast<VectorType>(C->getType())) {
+  if (C->getType()->isVectorTy()) {
     if (HasFn(C))
       return true;
     if (isa<ConstantAggregateZero>(C))
       return false;
-    if (isa<ScalableVectorType>(C->getType()))
-      return false;
 
-    for (unsigned i = 0, e = cast<FixedVectorType>(VTy)->getNumElements();
-         i != e; ++i) {
-      if (Constant *Elem = C->getAggregateElement(i))
-        if (HasFn(Elem))
-          return true;
-    }
+    return C->containsMatchingVectorElement(HasFn);
   }
 
   return false;
@@ -351,11 +344,22 @@ bool Constant::containsConstantExpression() const {
   if (isa<ConstantInt>(this) || isa<ConstantFP>(this))
     return false;
 
-  if (auto *VTy = dyn_cast<FixedVectorType>(getType())) {
-    for (unsigned i = 0, e = VTy->getNumElements(); i != e; ++i)
-      if (isa<ConstantExpr>(getAggregateElement(i)))
-        return true;
+  return containsMatchingVectorElement(IsaPred<ConstantExpr>);
+}
+
+bool Constant::containsMatchingVectorElement(
+    function_ref<bool(Constant *)> PredFn) const {
+  auto *FVTy = dyn_cast<FixedVectorType>(getType());
+  if (!FVTy)
+    return false;
+
+  unsigned NumElts = FVTy->getNumElements();
+  for (unsigned I = 0; I != NumElts; ++I) {
+    Constant *Elem = getAggregateElement(I);
+    if (Elem && PredFn(Elem))
+      return true;
   }
+
   return false;
 }
 

diff  --git a/llvm/lib/Transforms/InstCombine/InstCombineMulDivRem.cpp b/llvm/lib/Transforms/InstCombine/InstCombineMulDivRem.cpp
index 04512ab5d6715..39598dbbf6c67 100644
--- a/llvm/lib/Transforms/InstCombine/InstCombineMulDivRem.cpp
+++ b/llvm/lib/Transforms/InstCombine/InstCombineMulDivRem.cpp
@@ -1325,16 +1325,9 @@ Instruction *InstCombinerImpl::commonIDivRemTransforms(BinaryOperator &I) {
 
   // If any element of a constant divisor fixed width vector is zero or undef
   // the behavior is undefined and we can fold the whole op to poison.
-  auto *Op1C = dyn_cast<Constant>(Op1);
-  Type *Ty = I.getType();
-  auto *VTy = dyn_cast<FixedVectorType>(Ty);
-  if (Op1C && VTy) {
-    unsigned NumElts = VTy->getNumElements();
-    for (unsigned i = 0; i != NumElts; ++i) {
-      Constant *Elt = Op1C->getAggregateElement(i);
-      if (Elt && (Elt->isNullValue() || isa<UndefValue>(Elt)))
-        return replaceInstUsesWith(I, PoisonValue::get(Ty));
-    }
+  if (match(Op1, m_ContainsMatchingVectorElement(
+                     m_CombineOr(m_Zero(), m_UndefValue())))) {
+    return replaceInstUsesWith(I, PoisonValue::get(I.getType()));
   }
 
   if (Instruction *Phi = foldBinopWithPhiOperands(I))

diff  --git a/llvm/lib/Transforms/InstCombine/InstructionCombining.cpp b/llvm/lib/Transforms/InstCombine/InstructionCombining.cpp
index dc20d794b9dea..6d901e835a1bd 100644
--- a/llvm/lib/Transforms/InstCombine/InstructionCombining.cpp
+++ b/llvm/lib/Transforms/InstCombine/InstructionCombining.cpp
@@ -5516,15 +5516,10 @@ Instruction *InstCombinerImpl::visitFreeze(FreezeInst &I) {
     auto *VTy = dyn_cast<FixedVectorType>(Ty);
     if (!VTy)
       return nullptr;
-    unsigned NumElts = VTy->getNumElements();
-    Constant *BestValue = Constant::getNullValue(VTy->getScalarType());
-    for (unsigned i = 0; i != NumElts; ++i) {
-      Constant *EltC = C->getAggregateElement(i);
-      if (EltC && !match(EltC, m_Undef())) {
-        BestValue = EltC;
-        break;
-      }
-    }
+    Constant *BestValue;
+    if (!match(C, m_ContainsMatchingVectorElement(m_CombineAnd(
+                      m_Unless(m_Undef()), m_Constant(BestValue)))))
+      BestValue = Constant::getNullValue(VTy->getScalarType());
     return Constant::replaceUndefsWith(C, BestValue);
   };
 


        


More information about the llvm-commits mailing list