[llvm] [KnownFPClass] Add neg_square and fma_neg_square (PR #227248)
via llvm-commits
llvm-commits at lists.llvm.org
Wed Sep 30 02:54:07 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-llvm-analysis
Author: smv (s-mv)
<details>
<summary>Changes</summary>
Add KnownFPClass::neg_square and fma_neg_square and use them in computeKnownFPClass in ValueTracking and GISelValueTracking to recognise `fmul (fneg x), x`, and `fma/fmuladd (fneg x), x, y`.
+0 case is only excluded when input and output denormal modes can't flush to +0.
Part of #<!-- -->226728.
---
Full diff: https://github.com/llvm/llvm-project/pull/227248.diff
5 Files Affected:
- (modified) llvm/include/llvm/Support/KnownFPClass.h (+26)
- (modified) llvm/lib/Analysis/ValueTracking.cpp (+40-6)
- (modified) llvm/lib/CodeGen/GlobalISel/GISelValueTracking.cpp (+65-27)
- (modified) llvm/lib/Support/KnownFPClass.cpp (+14)
- (added) llvm/test/Transforms/InstCombine/known-fpclass-neg-square.ll (+167)
``````````diff
diff --git a/llvm/include/llvm/Support/KnownFPClass.h b/llvm/include/llvm/Support/KnownFPClass.h
index 22097a70a361b..ccd225df67326 100644
--- a/llvm/include/llvm/Support/KnownFPClass.h
+++ b/llvm/include/llvm/Support/KnownFPClass.h
@@ -287,6 +287,27 @@ struct KnownFPClass {
return Known;
}
+ // Special cases of fmul -x, x and fmul x, -x.
+ static KnownFPClass
+ neg_square(const KnownFPClass &Src,
+ DenormalMode Mode = DenormalMode::getDynamic()) {
+ KnownFPClass Known = fmul(fneg(Src), Src, Mode);
+
+ // -X * X is always negative, zero, or a NaN.
+ Known.knownNot(fcPosSubnormal | fcPosNormal | fcPosInf);
+
+ // Zero results are -0 unless a denormal is flushed to +0.
+ if ((Mode.Input == DenormalMode::IEEE ||
+ Mode.Input == DenormalMode::PreserveSign) &&
+ (Mode.Output == DenormalMode::IEEE ||
+ Mode.Output == DenormalMode::PreserveSign)) {
+ Known.knownNot(fcPosZero);
+ }
+
+ Known.propagateNonNaN(Src);
+ return Known;
+ }
+
LLVM_ABI static KnownFPClass
fmul(const KnownFPClass &LHS, const APFloat &RHS,
DenormalMode Mode = DenormalMode::getDynamic());
@@ -322,6 +343,11 @@ struct KnownFPClass {
fma_square(const KnownFPClass &Squared, const KnownFPClass &Addend,
DenormalMode Mode = DenormalMode::getDynamic());
+ /// Report known values for fma (-x, x, addend) and fma (x, -x, addend)
+ LLVM_ABI static KnownFPClass
+ fma_neg_square(const KnownFPClass &Squared, const KnownFPClass &Addend,
+ DenormalMode Mode = DenormalMode::getDynamic());
+
/// Propagate known class for sqrt
LLVM_ABI static KnownFPClass
sqrt(const KnownFPClass &Src, DenormalMode Mode = DenormalMode::getDynamic());
diff --git a/llvm/lib/Analysis/ValueTracking.cpp b/llvm/lib/Analysis/ValueTracking.cpp
index dd5f6fc7ad16a..3f6ebae3ec7a5 100644
--- a/llvm/lib/Analysis/ValueTracking.cpp
+++ b/llvm/lib/Analysis/ValueTracking.cpp
@@ -5402,16 +5402,32 @@ void computeKnownFPClass(const Value *V, const APInt &DemandedElts,
}
case Intrinsic::fma:
case Intrinsic::fmuladd: {
- if ((InterestedClasses & fcNegative) == fcNone)
+ Value *A0 = II->getArgOperand(0), *A1 = II->getArgOperand(1);
+ Value *NegSrc = nullptr;
+ if (match(A0, m_FNeg(m_Specific(A1))))
+ NegSrc = A1;
+ else if (match(A1, m_FNeg(m_Specific(A0))))
+ NegSrc = A0;
+
+ // Result is non-positive or NaN for -x * x and x * -x.
+ if (NegSrc && (InterestedClasses & (fcPositive | fcNan)) == fcNone) {
break;
+ }
+
+ // For x * x, result is positive or NaN.
+ if (!NegSrc && (InterestedClasses & fcNegative) == fcNone) {
+ break;
+ }
// FIXME: This should check isGuaranteedNotToBeUndef
- if (II->getArgOperand(0) == II->getArgOperand(1)) {
+ if (A0 == A1 || NegSrc) {
+ const Value *Src = NegSrc ? NegSrc : A0;
+
KnownFPClass KnownSrc, KnownAddend;
computeKnownFPClass(II->getArgOperand(2), DemandedElts,
InterestedClasses, KnownAddend, Q, Depth + 1);
- computeKnownFPClass(II->getArgOperand(0), DemandedElts,
- InterestedClasses, KnownSrc, Q, Depth + 1);
+ computeKnownFPClass(Src, DemandedElts, InterestedClasses, KnownSrc, Q,
+ Depth + 1);
const Function *F = II->getFunction();
const fltSemantics &FltSem =
@@ -5423,13 +5439,14 @@ void computeKnownFPClass(const Value *V, const APInt &DemandedElts,
KnownSrc.knownNot(fcNan);
KnownAddend.knownNot(fcNan);
}
-
if (KnownNotFromFlags & fcInf) {
KnownSrc.knownNot(fcInf);
KnownAddend.knownNot(fcInf);
}
- Known = KnownFPClass::fma_square(KnownSrc, KnownAddend, Mode);
+ Known = NegSrc
+ ? KnownFPClass::fma_neg_square(KnownSrc, KnownAddend, Mode)
+ : KnownFPClass::fma_square(KnownSrc, KnownAddend, Mode);
break;
}
@@ -6035,6 +6052,23 @@ void computeKnownFPClass(const Value *V, const APInt &DemandedElts,
break;
}
+ // Handling the -X * X and X * -X cases
+ Value *Src = nullptr;
+
+ if (match(LHS, m_FNeg(m_Specific(RHS)))) {
+ Src = RHS;
+ } else if (match(RHS, m_FNeg(m_Specific(LHS)))) {
+ Src = LHS;
+ }
+
+ if (Src) {
+ KnownFPClass KnownSrc;
+ computeKnownFPClass(Src, DemandedElts, fcAllFlags, KnownSrc, Q,
+ Depth + 1);
+ Known = KnownFPClass::neg_square(KnownSrc, Mode);
+ break;
+ }
+
KnownFPClass KnownLHS, KnownRHS;
const APFloat *CRHS;
diff --git a/llvm/lib/CodeGen/GlobalISel/GISelValueTracking.cpp b/llvm/lib/CodeGen/GlobalISel/GISelValueTracking.cpp
index 71b68e223959e..507ac2c4875ad 100644
--- a/llvm/lib/CodeGen/GlobalISel/GISelValueTracking.cpp
+++ b/llvm/lib/CodeGen/GlobalISel/GISelValueTracking.cpp
@@ -1416,49 +1416,72 @@ void GISelValueTracking::computeKnownFPClass(Register R,
case TargetOpcode::G_FMA:
case TargetOpcode::G_STRICT_FMA:
case TargetOpcode::G_FMAD: {
- if ((InterestedClasses & fcNegative) == fcNone)
- break;
-
Register A = MI.getOperand(1).getReg();
Register B = MI.getOperand(2).getReg();
Register C = MI.getOperand(3).getReg();
+ // Match (-x) * x or x * (-x).
+ Register NegSrc;
+
+ if (mi_match(A, MRI, m_GFNeg(m_SpecificReg(B))))
+ NegSrc = B;
+ else if (mi_match(B, MRI, m_GFNeg(m_SpecificReg(A))))
+ NegSrc = A;
+
+ // Result is non-positive or NaN.
+ if (NegSrc.isValid() &&
+ (InterestedClasses & (fcPositive | fcNan)) == fcNone) {
+ break;
+ }
+
+ // For x * x, result is positive or NaN.
+ if (!NegSrc.isValid() && (InterestedClasses & fcNegative) == fcNone) {
+ break;
+ }
+
DenormalMode Mode =
MF->getDenormalMode(getFltSemanticForLLT(DstTy.getScalarType()));
- if (A == B && isGuaranteedNotToBeUndef(A, MRI, Depth + 1)) {
- // x * x + y
+ if ((A == B && isGuaranteedNotToBeUndef(A, MRI, Depth + 1)) ||
+ (NegSrc.isValid() &&
+ isGuaranteedNotToBeUndef(NegSrc, MRI, Depth + 1))) {
+ Register Src = NegSrc.isValid() ? NegSrc : A;
KnownFPClass KnownSrc, KnownAddend;
+
computeKnownFPClass(C, DemandedElts, InterestedClasses, KnownAddend,
Depth + 1);
- computeKnownFPClass(A, DemandedElts, InterestedClasses, KnownSrc,
+ computeKnownFPClass(Src, DemandedElts, InterestedClasses, KnownSrc,
Depth + 1);
+
if (KnownNotFromFlags) {
KnownSrc.knownNot(KnownNotFromFlags);
KnownAddend.knownNot(KnownNotFromFlags);
}
- Known = KnownFPClass::fma_square(KnownSrc, KnownAddend, Mode);
- } else {
- KnownFPClass KnownSrc[3];
- computeKnownFPClass(A, DemandedElts, InterestedClasses, KnownSrc[0],
- Depth + 1);
- if (KnownSrc[0].isUnknown())
- break;
- computeKnownFPClass(B, DemandedElts, InterestedClasses, KnownSrc[1],
- Depth + 1);
- if (KnownSrc[1].isUnknown())
- break;
- computeKnownFPClass(C, DemandedElts, InterestedClasses, KnownSrc[2],
- Depth + 1);
- if (KnownSrc[2].isUnknown())
- break;
- if (KnownNotFromFlags) {
- KnownSrc[0].knownNot(KnownNotFromFlags);
- KnownSrc[1].knownNot(KnownNotFromFlags);
- KnownSrc[2].knownNot(KnownNotFromFlags);
- }
- Known = KnownFPClass::fma(KnownSrc[0], KnownSrc[1], KnownSrc[2], Mode);
+
+ Known = NegSrc.isValid()
+ ? KnownFPClass::fma_neg_square(KnownSrc, KnownAddend, Mode)
+ : KnownFPClass::fma_square(KnownSrc, KnownAddend, Mode);
+ break;
+ }
+
+ KnownFPClass KnownSrc[3];
+ computeKnownFPClass(A, DemandedElts, InterestedClasses, KnownSrc[0],
+ Depth + 1);
+ if (KnownSrc[0].isUnknown())
+ break;
+ computeKnownFPClass(B, DemandedElts, InterestedClasses, KnownSrc[1],
+ Depth + 1);
+ if (KnownSrc[1].isUnknown())
+ break;
+ computeKnownFPClass(C, DemandedElts, InterestedClasses, KnownSrc[2],
+ Depth + 1);
+ if (KnownSrc[2].isUnknown())
+ break;
+ if (KnownNotFromFlags) {
+ for (KnownFPClass &K : KnownSrc)
+ K.knownNot(KnownNotFromFlags);
}
+ Known = KnownFPClass::fma(KnownSrc[0], KnownSrc[1], KnownSrc[2], Mode);
break;
}
case TargetOpcode::G_FSQRT:
@@ -1896,6 +1919,21 @@ void GISelValueTracking::computeKnownFPClass(Register R,
computeKnownFPClass(LHS, DemandedElts, fcAllFlags, KnownSrc, Depth + 1);
Known = KnownFPClass::square(KnownSrc, Mode);
} else {
+ // -X * X and X * -X cases.
+ Register Src;
+ if (mi_match(LHS, MRI, m_GFNeg(m_SpecificReg(RHS)))) {
+ Src = RHS;
+ } else if (mi_match(RHS, MRI, m_GFNeg(m_SpecificReg(LHS)))) {
+ Src = LHS;
+ }
+
+ if (Src.isValid() && isGuaranteedNotToBeUndef(Src, MRI, Depth + 1)) {
+ KnownFPClass KnownSrc;
+ computeKnownFPClass(Src, DemandedElts, fcAllFlags, KnownSrc, Depth + 1);
+ Known = KnownFPClass::neg_square(KnownSrc, Mode);
+ break;
+ }
+
// If RHS is a scalar constant, use the more precise APFloat overload.
auto RHSCst = GFConstant::getConstant(RHS, MRI);
if (RHSCst && RHSCst->getKind() == GFConstant::GFConstantKind::Scalar) {
diff --git a/llvm/lib/Support/KnownFPClass.cpp b/llvm/lib/Support/KnownFPClass.cpp
index 1451a407983ff..58cc1b5912a9e 100644
--- a/llvm/lib/Support/KnownFPClass.cpp
+++ b/llvm/lib/Support/KnownFPClass.cpp
@@ -640,6 +640,20 @@ KnownFPClass KnownFPClass::fma_square(const KnownFPClass &KnownSquared,
return Known;
}
+KnownFPClass KnownFPClass::fma_neg_square(const KnownFPClass &KnownSquared,
+ const KnownFPClass &KnownAddend,
+ DenormalMode Mode) {
+ KnownFPClass NegSquared = neg_square(KnownSquared, Mode);
+ KnownFPClass Known = fadd_impl(NegSquared, KnownAddend, Mode);
+
+ if (KnownAddend.isKnownNever(fcPosInf | fcNan) &&
+ NegSquared.isKnownNever(fcNan))
+ Known.knownNot(fcNan);
+
+ Known.propagateNonSNaN(KnownSquared, KnownAddend);
+ return Known;
+}
+
KnownFPClass KnownFPClass::exp(const KnownFPClass &KnownSrc) {
KnownFPClass Known;
Known.knownNot(fcNegative);
diff --git a/llvm/test/Transforms/InstCombine/known-fpclass-neg-square.ll b/llvm/test/Transforms/InstCombine/known-fpclass-neg-square.ll
new file mode 100644
index 0000000000000..0d3de1cafb55c
--- /dev/null
+++ b/llvm/test/Transforms/InstCombine/known-fpclass-neg-square.ll
@@ -0,0 +1,167 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --version 6
+; RUN: opt -S -passes=instcombine < %s | FileCheck %s
+
+declare float @llvm.fma.f32(float, float, float)
+declare float @llvm.fabs.f32(float)
+declare float @llvm.fmuladd.f32(float, float, float)
+declare i1 @llvm.is.fpclass.f32(float, i32)
+
+; (-x) * x is never a positive non-zero value, so this must fold
+define i1 @fmul_neg_not_pos(float %x) {
+; CHECK-LABEL: define i1 @fmul_neg_not_pos(
+; CHECK-SAME: float [[X:%.*]]) {
+; CHECK-NEXT: ret i1 false
+;
+ %neg = fneg float %x ; -x
+ %res = fmul float %neg, %x ; (-x) * x
+ %r = call i1 @llvm.is.fpclass.f32(float %res, i32 896)
+ ret i1 %r
+}
+
+; test for x * (-x)
+define i1 @fmul_neg_not_pos_commutative(float %x) {
+; CHECK-LABEL: define i1 @fmul_neg_not_pos_commutative(
+; CHECK-SAME: float [[X:%.*]]) {
+; CHECK-NEXT: ret i1 false
+;
+ %neg = fneg float %x ; -x
+ %res = fmul float %x, %neg ; x * (-x)
+ %r = call i1 @llvm.is.fpclass.f32(float %res, i32 896)
+ ret i1 %r
+}
+
+; flush to positive zero case should not fold
+define i1 @fmul_neg_pzero(float %x) #0 {
+; CHECK-LABEL: define i1 @fmul_neg_pzero(
+; CHECK-SAME: float [[X:%.*]]) #[[ATTR1:[0-9]+]] {
+; CHECK-NEXT: [[NEG:%.*]] = fneg float [[X]]
+; CHECK-NEXT: [[RES:%.*]] = fmul float [[X]], [[NEG]]
+; CHECK-NEXT: [[R:%.*]] = call i1 @llvm.is.fpclass.f32(float [[RES]], /* (pzero) */ i32 64)
+; CHECK-NEXT: ret i1 [[R]]
+;
+ %neg = fneg float %x
+ %res = fmul float %neg, %x
+ %r = call i1 @llvm.is.fpclass.f32(float %res, i32 64)
+ ret i1 %r
+}
+
+; with IEEE denormals this must fold
+define i1 @fmul_neg_pzero_ieee(float %x) #1 {
+; CHECK-LABEL: define i1 @fmul_neg_pzero_ieee(
+; CHECK-SAME: float [[X:%.*]]) #[[ATTR2:[0-9]+]] {
+; CHECK-NEXT: ret i1 false
+;
+ %neg = fneg float %x
+ %res = fmul float %neg, %x
+ %r = call i1 @llvm.is.fpclass.f32(float %res, i32 64)
+ ret i1 %r
+}
+
+; should still fold with preservesign
+define i1 @fmul_neg_pzero_preservesign(float %x) #2 {
+; CHECK-LABEL: define i1 @fmul_neg_pzero_preservesign(
+; CHECK-SAME: float [[X:%.*]]) #[[ATTR3:[0-9]+]] {
+; CHECK-NEXT: ret i1 false
+;
+ %neg = fneg float %x
+ %res = fmul float %neg, %x
+ %r = call i1 @llvm.is.fpclass.f32(float %res, i32 64)
+ ret i1 %r
+}
+
+; shouldn't fold
+define i1 @fmul_neg_pzero_dynamic(float %x) #3 {
+; CHECK-LABEL: define i1 @fmul_neg_pzero_dynamic(
+; CHECK-SAME: float [[X:%.*]]) #[[ATTR4:[0-9]+]] {
+; CHECK-NEXT: [[NEG:%.*]] = fneg float [[X]]
+; CHECK-NEXT: [[RES:%.*]] = fmul float [[X]], [[NEG]]
+; CHECK-NEXT: [[R:%.*]] = call i1 @llvm.is.fpclass.f32(float [[RES]], /* (pzero) */ i32 64)
+; CHECK-NEXT: ret i1 [[R]]
+;
+ %neg = fneg float %x
+ %res = fmul float %neg, %x
+ %r = call i1 @llvm.is.fpclass.f32(float %res, i32 64)
+ ret i1 %r
+}
+
+; only output side can flush to +0 so it must not fold
+define i1 @fmul_neg_pzero_out_pzero(float %x) #4 {
+; CHECK-LABEL: define i1 @fmul_neg_pzero_out_pzero(
+; CHECK-SAME: float [[X:%.*]]) #[[ATTR5:[0-9]+]] {
+; CHECK-NEXT: [[NEG:%.*]] = fneg float [[X]]
+; CHECK-NEXT: [[RES:%.*]] = fmul float [[X]], [[NEG]]
+; CHECK-NEXT: [[R:%.*]] = call i1 @llvm.is.fpclass.f32(float [[RES]], /* (pzero) */ i32 64)
+; CHECK-NEXT: ret i1 [[R]]
+;
+ %neg = fneg float %x
+ %res = fmul float %neg, %x
+ %r = call i1 @llvm.is.fpclass.f32(float %res, i32 64)
+ ret i1 %r
+}
+
+; fma of the form -x, x, -y
+define i1 @fma_negx_x_negy(float %x, float %y) {
+; CHECK-LABEL: define i1 @fma_negx_x_negy(
+; CHECK-SAME: float [[X:%.*]], float [[Y:%.*]]) {
+; CHECK-NEXT: ret i1 false
+;
+ %neg1 = fneg float %x ; -x
+ %abs = call float @llvm.fabs.f32(float %y)
+ %neg2 = fneg float %abs ; -y
+ %res = call float @llvm.fma.f32(float %neg1, float %x, float %neg2)
+ %r = call i1 @llvm.is.fpclass.f32(float %res, i32 896)
+ ret i1 %r
+}
+
+; fma contracted from -(x * x) + y shouldn't fold
+define i1 @fma_y_minus_neg_x_squared(float %x, float %y) {
+; CHECK-LABEL: define i1 @fma_y_minus_neg_x_squared(
+; CHECK-SAME: float [[X:%.*]], float [[Y:%.*]]) {
+; CHECK-NEXT: [[NEG:%.*]] = fneg float [[X]]
+; CHECK-NEXT: [[ABS:%.*]] = call float @llvm.fabs.f32(float [[Y]])
+; CHECK-NEXT: [[RES:%.*]] = call float @llvm.fmuladd.f32(float [[NEG]], float [[X]], float [[ABS]])
+; CHECK-NEXT: [[R:%.*]] = fcmp uno float [[RES]], 0.000000e+00
+; CHECK-NEXT: ret i1 [[R]]
+;
+ %neg = fneg float %x
+ %abs = call float @llvm.fabs.f32(float %y)
+ %res = call float @llvm.fmuladd.f32(float %neg, float %x, float %abs)
+ %r = call i1 @llvm.is.fpclass.f32(float %res, i32 3)
+ ret i1 %r
+}
+
+; must not fold due to possibility of nan
+define i1 @fma_neg_x_x_pinf_y(float nofpclass(nan) %x, float nofpclass(nan) %y) {
+; CHECK-LABEL: define i1 @fma_neg_x_x_pinf_y(
+; CHECK-SAME: float nofpclass(nan) [[X:%.*]], float nofpclass(nan) [[Y:%.*]]) {
+; CHECK-NEXT: [[NEG:%.*]] = fneg float [[X]]
+; CHECK-NEXT: [[RES:%.*]] = call float @llvm.fma.f32(float [[NEG]], float [[X]], float [[Y]])
+; CHECK-NEXT: [[R:%.*]] = fcmp uno float [[RES]], 0.000000e+00
+; CHECK-NEXT: ret i1 [[R]]
+;
+ %neg = fneg float %x
+ %res = call float @llvm.fma.f32(float %neg, float %x, float %y)
+ %r = call i1 @llvm.is.fpclass.f32(float %res, i32 3)
+ ret i1 %r
+}
+
+; must not fold since the addend may be of any class
+define i1 @fma_neg_x_x_unknown_y(float %x, float %y) {
+; CHECK-LABEL: define i1 @fma_neg_x_x_unknown_y(
+; CHECK-SAME: float [[X:%.*]], float [[Y:%.*]]) {
+; CHECK-NEXT: [[NEG:%.*]] = fneg float [[X]]
+; CHECK-NEXT: [[RES:%.*]] = call float @llvm.fma.f32(float [[NEG]], float [[X]], float [[Y]])
+; CHECK-NEXT: [[R:%.*]] = fcmp ogt float [[RES]], 0.000000e+00
+; CHECK-NEXT: ret i1 [[R]]
+;
+ %neg = fneg float %x
+ %res = call float @llvm.fma.f32(float %neg, float %x, float %y)
+ %r = call i1 @llvm.is.fpclass.f32(float %res, i32 896)
+ ret i1 %r
+}
+
+attributes #0 = { denormal_fpenv(positivezero) }
+attributes #1 = { denormal_fpenv(ieee) }
+attributes #2 = { denormal_fpenv(preservesign) }
+attributes #3 = { denormal_fpenv(dynamic) }
+attributes #4 = { denormal_fpenv(positivezero|ieee) }
``````````
</details>
https://github.com/llvm/llvm-project/pull/227248
More information about the llvm-commits
mailing list