@@ -208,7 +208,7 @@ class AMDGPUCodeGenPrepareImpl
208208
209209 bool canWidenScalarExtLoad (LoadInst &I) const ;
210210
211- Value *matchFractPat (IntrinsicInst &I );
211+ Value *matchFractPat (Value &V );
212212 Value *applyFractPat (IRBuilder<> &Builder, Value *FractArg);
213213
214214 bool canOptimizeWithRsq (FastMathFlags DivFMF, FastMathFlags SqrtFMF) const ;
@@ -1595,10 +1595,10 @@ bool AMDGPUCodeGenPrepareImpl::visitSelectInst(SelectInst &I) {
15951595 Value *TrueVal = I.getTrueValue ();
15961596 Value *FalseVal = I.getFalseValue ();
15971597 Value *CmpVal;
1598- CmpPredicate Pred ;
1598+ CmpPredicate IsNanPred ;
15991599
16001600 // Match fract pattern with nan check.
1601- if (!match (Cond, m_FCmp (Pred , m_Value (CmpVal), m_NonNaN ())))
1601+ if (!match (Cond, m_FCmp (IsNanPred , m_Value (CmpVal), m_NonNaN ())))
16021602 return false ;
16031603
16041604 FPMathOperator *FPOp = dyn_cast<FPMathOperator>(&I);
@@ -1608,18 +1608,44 @@ bool AMDGPUCodeGenPrepareImpl::visitSelectInst(SelectInst &I) {
16081608 IRBuilder<> Builder (&I);
16091609 Builder.setFastMathFlags (FPOp->getFastMathFlags ());
16101610
1611- auto *IITrue = dyn_cast<IntrinsicInst>(TrueVal);
1612- auto *IIFalse = dyn_cast<IntrinsicInst>(FalseVal);
1613-
16141611 Value *Fract = nullptr ;
1615- if (Pred == FCmpInst::FCMP_UNO && TrueVal == CmpVal && IIFalse &&
1616- CmpVal == matchFractPat (*IIFalse )) {
1612+ if (IsNanPred == FCmpInst::FCMP_UNO && TrueVal == CmpVal &&
1613+ CmpVal == matchFractPat (*FalseVal )) {
16171614 // isnan(x) ? x : fract(x)
16181615 Fract = applyFractPat (Builder, CmpVal);
1619- } else if (Pred == FCmpInst::FCMP_ORD && FalseVal == CmpVal && IITrue &&
1620- CmpVal == matchFractPat (*IITrue)) {
1621- // !isnan(x) ? fract(x) : x
1622- Fract = applyFractPat (Builder, CmpVal);
1616+ } else if (IsNanPred == FCmpInst::FCMP_ORD && FalseVal == CmpVal) {
1617+ if (CmpVal == matchFractPat (*TrueVal)) {
1618+ // !isnan(x) ? fract(x) : x
1619+ Fract = applyFractPat (Builder, CmpVal);
1620+ } else {
1621+ // Match an intermediate clamp infinity to 0 pattern. i.e.
1622+ // !isnan(x) ? (!isinf(x) ? fract(x) : 0.0) : x
1623+ CmpPredicate PredInf;
1624+ Value *IfNotInf;
1625+
1626+ if (!match (TrueVal, m_Select (m_FCmp (PredInf, m_FAbs (m_Specific (CmpVal)),
1627+ m_PosInf ()),
1628+ m_Value (IfNotInf), m_PosZeroFP ())) ||
1629+ PredInf != FCmpInst::FCMP_UNE || CmpVal != matchFractPat (*IfNotInf))
1630+ return false ;
1631+
1632+ SelectInst *ClampInfSelect = cast<SelectInst>(TrueVal);
1633+
1634+ // Insert before the fabs
1635+ Value *InsertPt =
1636+ cast<Instruction>(ClampInfSelect->getCondition ())->getOperand (0 );
1637+
1638+ Builder.SetInsertPoint (cast<Instruction>(InsertPt));
1639+ Value *NewFract = applyFractPat (Builder, CmpVal);
1640+ NewFract->takeName (TrueVal);
1641+
1642+ // Thread the new fract into the inf clamping sequence.
1643+ DeadVals.push_back (ClampInfSelect->getOperand (1 ));
1644+ ClampInfSelect->setOperand (1 , NewFract);
1645+
1646+ // The outer select nan handling is also absorbed into the fract.
1647+ Fract = ClampInfSelect;
1648+ }
16231649 } else
16241650 return false ;
16251651
@@ -2029,24 +2055,28 @@ bool AMDGPUCodeGenPrepareImpl::visitIntrinsicInst(IntrinsicInst &I) {
20292055// /
20302056// / If fract is a useful instruction for the subtarget. Does not account for the
20312057// / nan handling; the instruction has a nan check on the input value.
2032- Value *AMDGPUCodeGenPrepareImpl::matchFractPat (IntrinsicInst &I ) {
2058+ Value *AMDGPUCodeGenPrepareImpl::matchFractPat (Value &V ) {
20332059 if (ST .hasFractBug ())
20342060 return nullptr ;
20352061
2036- Intrinsic::ID IID = I.getIntrinsicID ();
2062+ IntrinsicInst *II = dyn_cast<IntrinsicInst>(&V);
2063+ if (!II )
2064+ return nullptr ;
2065+
2066+ Intrinsic::ID IID = II ->getIntrinsicID ();
20372067
20382068 // The value is only used in contexts where we know the input isn't a nan, so
20392069 // any of the fmin variants are fine.
20402070 if (IID != Intrinsic::minnum && IID != Intrinsic::minimum &&
20412071 IID != Intrinsic::minimumnum)
20422072 return nullptr ;
20432073
2044- Type *Ty = I .getType ();
2074+ Type *Ty = V .getType ();
20452075 if (!isLegalFloatingTy (Ty->getScalarType ()))
20462076 return nullptr ;
20472077
2048- Value *Arg0 = I. getArgOperand (0 );
2049- Value *Arg1 = I. getArgOperand (1 );
2078+ Value *Arg0 = II -> getArgOperand (0 );
2079+ Value *Arg1 = II -> getArgOperand (1 );
20502080
20512081 const APFloat *C;
20522082 if (!match (Arg1, m_APFloatAllowPoison (C)))
0 commit comments