[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