[llvm] [ARM] Have SelectionDAG optimize muls, not ISelDAGtoDAG (PR #195334)

via llvm-commits llvm-commits at lists.llvm.org
Fri May 1 12:24:23 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-backend-arm

Author: LumioseSil (LumioseSil)

<details>
<summary>Changes</summary>

I had to swap so shifts would be on the right side for better folding.

Made that change in AArch64 too for parity.

---

Patch is 84.03 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/195334.diff


19 Files Affected:

- (modified) llvm/lib/Target/AArch64/AArch64ISelLowering.cpp (+10-10) 
- (modified) llvm/lib/Target/ARM/ARMISelDAGToDAG.cpp (-46) 
- (modified) llvm/lib/Target/ARM/ARMISelLowering.cpp (+202-68) 
- (modified) llvm/test/CodeGen/AArch64/load-insert-zero.ll (+2-2) 
- (modified) llvm/test/CodeGen/AArch64/sme2-intrinsics-int-dots.ll (+14-14) 
- (modified) llvm/test/CodeGen/AArch64/sme2-intrinsics-vdot.ll (+8-8) 
- (modified) llvm/test/CodeGen/ARM/2013-05-07-ByteLoadSameAddress.ll (+1-1) 
- (modified) llvm/test/CodeGen/ARM/addimm-mulimm.ll (+134-90) 
- (modified) llvm/test/CodeGen/ARM/funnel-shift.ll (+12-8) 
- (modified) llvm/test/CodeGen/ARM/memset-inline.ll (+155-34) 
- (modified) llvm/test/CodeGen/ARM/mul_const.ll (+4-4) 
- (modified) llvm/test/CodeGen/ARM/popcnt.ll (+28-31) 
- (modified) llvm/test/CodeGen/ARM/select-imm.ll (+8-9) 
- (modified) llvm/test/CodeGen/ARM/srem-seteq-illegal-types.ll (+14-12) 
- (modified) llvm/test/CodeGen/ARM/urem-seteq-illegal-types.ll (+13-13) 
- (modified) llvm/test/CodeGen/Thumb2/mve-gather-scatter-optimisation.ll (+24-24) 
- (modified) llvm/test/CodeGen/Thumb2/mve-memtp-branch.ll (+14-13) 
- (modified) llvm/test/CodeGen/Thumb2/mve-postinc-dct.ll (+55-55) 
- (modified) llvm/test/CodeGen/Thumb2/urem-seteq-illegal-types.ll (+2-2) 


``````````diff
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index b23bbd72341778..572b3a3f0d51a2 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -20520,13 +20520,13 @@ static SDValue performMulCombine(SDNode *N, SelectionDAG &DAG,
   };
 
   if (ConstValue.isNonNegative()) {
-    // (mul x, (2^N + 1) * 2^M) => (shl (add (shl x, N), x), M)
+    // (mul x, (2^N + 1) * 2^M) => (shl (add x, (shl x, N)), M)
     // (mul x, 2^N - 1) => (sub (shl x, N), x)
     // (mul x, (2^(N-M) - 1) * 2^M) => (sub (shl x, N), (shl x, M))
     // (mul x, (2^M + 1) * (2^N + 1))
-    //     => MV = (add (shl x, M), x); (add (shl MV, N), MV)
+    //     => MV = (add x, (shl x, M)); (add MV, (shl MV, N))
     // (mul x, (2^M + 1) * 2^N + 1))
-    //     =>  MV = add (shl x, M), x); add (shl MV, N), x)
+    //     =>  MV = (add x, (shl x, M)); (add x, (shl MV, N)))
     // (mul x, 1 - (1 - 2^M) * 2^N))
     //     =>  MV = sub (x - (shl x, M)); sub (x - (shl MV, N))
     APInt SCVMinus1 = ShiftedConstValue - 1;
@@ -20535,7 +20535,7 @@ static SDValue performMulCombine(SDNode *N, SelectionDAG &DAG,
     APInt CVM, CVN;
     if (SCVMinus1.isPowerOf2()) {
       ShiftAmt = SCVMinus1.logBase2();
-      return Shl(Add(Shl(N0, ShiftAmt), N0), TrailingZeroes);
+      return Shl(Add(N0, Shl(N0, ShiftAmt)), TrailingZeroes);
     } else if (CVPlus1.isPowerOf2()) {
       ShiftAmt = CVPlus1.logBase2();
       return Sub(Shl(N0, ShiftAmt), N0);
@@ -20551,8 +20551,8 @@ static SDValue performMulCombine(SDNode *N, SelectionDAG &DAG,
       unsigned ShiftN1 = CVNMinus1.logBase2();
       // ALULSLFast implicate that Shifts <= 4 places are fast
       if (ShiftM1 <= 4 && ShiftN1 <= 4) {
-        SDValue MVal = Add(Shl(N0, ShiftM1), N0);
-        return Add(Shl(MVal, ShiftN1), MVal);
+        SDValue MVal = Add(N0, Shl(N0, ShiftM1));
+        return Add(MVal, Shl(MVal, ShiftN1));
       }
     }
     if (Subtarget->hasALULSLFast() &&
@@ -20561,8 +20561,8 @@ static SDValue performMulCombine(SDNode *N, SelectionDAG &DAG,
       unsigned ShiftN = CVN.getZExtValue();
       // ALULSLFast implicate that Shifts <= 4 places are fast
       if (ShiftM <= 4 && ShiftN <= 4) {
-        SDValue MVal = Add(Shl(N0, CVM.getZExtValue()), N0);
-        return Add(Shl(MVal, CVN.getZExtValue()), N0);
+        SDValue MVal = Add(N0, Shl(N0, CVM.getZExtValue()));
+        return Add(N0, Shl(MVal, CVN.getZExtValue()));
       }
     }
 
@@ -20578,7 +20578,7 @@ static SDValue performMulCombine(SDNode *N, SelectionDAG &DAG,
     }
   } else {
     // (mul x, -(2^N - 1)) => (sub x, (shl x, N))
-    // (mul x, -(2^N + 1)) => - (add (shl x, N), x)
+    // (mul x, -(2^N + 1)) => - (add x, (shl x, N)))
     // (mul x, -(2^(N-M) - 1) * 2^M) => (sub (shl x, M), (shl x, N))
     APInt SCVPlus1 = -ShiftedConstValue + 1;
     APInt CVNegPlus1 = -ConstValue + 1;
@@ -20588,7 +20588,7 @@ static SDValue performMulCombine(SDNode *N, SelectionDAG &DAG,
       return Sub(N0, Shl(N0, ShiftAmt));
     } else if (CVNegMinus1.isPowerOf2()) {
       ShiftAmt = CVNegMinus1.logBase2();
-      return Negate(Add(Shl(N0, ShiftAmt), N0));
+      return Negate(Add(N0, Shl(N0, ShiftAmt)));
     } else if (SCVPlus1.isPowerOf2()) {
       ShiftAmt = SCVPlus1.logBase2() + TrailingZeroes;
       return Sub(Shl(N0, TrailingZeroes), Shl(N0, ShiftAmt));
diff --git a/llvm/lib/Target/ARM/ARMISelDAGToDAG.cpp b/llvm/lib/Target/ARM/ARMISelDAGToDAG.cpp
index 61b679d55fb47c..6d8a0276354cd4 100644
--- a/llvm/lib/Target/ARM/ARMISelDAGToDAG.cpp
+++ b/llvm/lib/Target/ARM/ARMISelDAGToDAG.cpp
@@ -3839,52 +3839,6 @@ void ARMDAGToDAGISel::Select(SDNode *N) {
     if (tryFMULFixed(N, dl))
       return;
     break;
-  case ISD::MUL:
-    if (Subtarget->isThumb1Only())
-      break;
-    if (ConstantSDNode *C = dyn_cast<ConstantSDNode>(N->getOperand(1))) {
-      unsigned RHSV = C->getZExtValue();
-      if (!RHSV) break;
-      if (isPowerOf2_32(RHSV-1)) {  // 2^n+1?
-        unsigned ShImm = Log2_32(RHSV-1);
-        if (ShImm >= 32)
-          break;
-        SDValue V = N->getOperand(0);
-        ShImm = ARM_AM::getSORegOpc(ARM_AM::lsl, ShImm);
-        SDValue ShImmOp = CurDAG->getTargetConstant(ShImm, dl, MVT::i32);
-        SDValue Reg0 = CurDAG->getRegister(0, MVT::i32);
-        if (Subtarget->isThumb()) {
-          SDValue Ops[] = { V, V, ShImmOp, getAL(CurDAG, dl), Reg0, Reg0 };
-          CurDAG->SelectNodeTo(N, ARM::t2ADDrs, MVT::i32, Ops);
-          return;
-        } else {
-          SDValue Ops[] = { V, V, Reg0, ShImmOp, getAL(CurDAG, dl), Reg0,
-                            Reg0 };
-          CurDAG->SelectNodeTo(N, ARM::ADDrsi, MVT::i32, Ops);
-          return;
-        }
-      }
-      if (isPowerOf2_32(RHSV+1)) {  // 2^n-1?
-        unsigned ShImm = Log2_32(RHSV+1);
-        if (ShImm >= 32)
-          break;
-        SDValue V = N->getOperand(0);
-        ShImm = ARM_AM::getSORegOpc(ARM_AM::lsl, ShImm);
-        SDValue ShImmOp = CurDAG->getTargetConstant(ShImm, dl, MVT::i32);
-        SDValue Reg0 = CurDAG->getRegister(0, MVT::i32);
-        if (Subtarget->isThumb()) {
-          SDValue Ops[] = { V, V, ShImmOp, getAL(CurDAG, dl), Reg0, Reg0 };
-          CurDAG->SelectNodeTo(N, ARM::t2RSBrs, MVT::i32, Ops);
-          return;
-        } else {
-          SDValue Ops[] = { V, V, Reg0, ShImmOp, getAL(CurDAG, dl), Reg0,
-                            Reg0 };
-          CurDAG->SelectNodeTo(N, ARM::RSBrsi, MVT::i32, Ops);
-          return;
-        }
-      }
-    }
-    break;
   case ISD::AND: {
     // Check for unsigned bitfield extract
     if (tryV6T2BitfieldExtractOp(N, false))
diff --git a/llvm/lib/Target/ARM/ARMISelLowering.cpp b/llvm/lib/Target/ARM/ARMISelLowering.cpp
index 71cc6cf8e1f820..e4ecc8642e75e7 100644
--- a/llvm/lib/Target/ARM/ARMISelLowering.cpp
+++ b/llvm/lib/Target/ARM/ARMISelLowering.cpp
@@ -14078,7 +14078,10 @@ static SDValue PerformSUBCombine(SDNode *N,
 static SDValue PerformVMULCombine(SDNode *N,
                                   TargetLowering::DAGCombinerInfo &DCI,
                                   const ARMSubtarget *Subtarget) {
-  if (!Subtarget->hasVMLxForwarding())
+
+  EVT VT = N->getValueType(0);
+  if (!Subtarget->hasVMLxForwarding() ||
+      (!VT.is64BitVector() && !VT.is128BitVector()))
     return SDValue();
 
   SelectionDAG &DAG = DCI.DAG;
@@ -14097,7 +14100,6 @@ static SDValue PerformVMULCombine(SDNode *N,
   if (N0 == N1)
     return SDValue();
 
-  EVT VT = N->getValueType(0);
   SDLoc DL(N);
   SDValue N00 = N0->getOperand(0);
   SDValue N01 = N0->getOperand(1);
@@ -14109,7 +14111,7 @@ static SDValue PerformVMULCombine(SDNode *N,
 static SDValue PerformMVEVMULLCombine(SDNode *N, SelectionDAG &DAG,
                                       const ARMSubtarget *Subtarget) {
   EVT VT = N->getValueType(0);
-  if (VT != MVT::v2i64)
+  if (!Subtarget->hasMVEIntegerOps() || VT != MVT::v2i64)
     return SDValue();
 
   SDValue N0 = N->getOperand(0);
@@ -14171,14 +14173,12 @@ static SDValue PerformMVEVMULLCombine(SDNode *N, SelectionDAG &DAG,
   return SDValue();
 }
 
-static SDValue PerformMULCombine(SDNode *N,
+static SDValue PerformMULCombine(SDNode *N, SelectionDAG &DAG,
                                  TargetLowering::DAGCombinerInfo &DCI,
                                  const ARMSubtarget *Subtarget) {
-  SelectionDAG &DAG = DCI.DAG;
-
-  EVT VT = N->getValueType(0);
-  if (Subtarget->hasMVEIntegerOps() && VT == MVT::v2i64)
-    return PerformMVEVMULLCombine(N, DAG, Subtarget);
+  
+  if (SDValue Val = PerformMVEVMULLCombine(N, DAG, Subtarget))
+    return Val;
 
   if (Subtarget->isThumb1Only())
     return SDValue();
@@ -14186,74 +14186,208 @@ static SDValue PerformMULCombine(SDNode *N,
   if (DCI.isBeforeLegalize() || DCI.isCalledByLegalizer())
     return SDValue();
 
-  if (VT.is64BitVector() || VT.is128BitVector())
-    return PerformVMULCombine(N, DCI, Subtarget);
-  if (VT != MVT::i32)
-    return SDValue();
+  if (SDValue Val = PerformVMULCombine(N, DCI, Subtarget))
+    return Val;
 
-  ConstantSDNode *C = dyn_cast<ConstantSDNode>(N->getOperand(1));
+  SDLoc DL(N);
+  EVT VT = N->getValueType(0);
+  SDValue N0 = N->getOperand(0);
+  SDValue N1 = N->getOperand(1);
+  SDValue MulOper;
+  unsigned AddSubOpc;
+
+  if (!Subtarget->isThumb()) {
+    auto IsAddSubWith1 = [&](SDValue V) -> bool {
+      AddSubOpc = V->getOpcode();
+      if ((AddSubOpc == ISD::ADD || AddSubOpc == ISD::SUB) && V->hasOneUse()) {
+        SDValue Opnd = V->getOperand(1);
+        MulOper = V->getOperand(0);
+        if (AddSubOpc == ISD::SUB)
+          std::swap(Opnd, MulOper);
+        if (auto C = dyn_cast<ConstantSDNode>(Opnd))
+          return C->isOne();
+      }
+      return false;
+    };
+
+    if (IsAddSubWith1(N0)) {
+      SDValue MulVal = DAG.getNode(ISD::MUL, DL, VT, N1, MulOper);
+      return DAG.getNode(AddSubOpc, DL, VT, N1, MulVal);
+    }
+
+    if (IsAddSubWith1(N1)) {
+      SDValue MulVal = DAG.getNode(ISD::MUL, DL, VT, N0, MulOper);
+      return DAG.getNode(AddSubOpc, DL, VT, N0, MulVal);
+    }
+  }
+
+  // The below optimizations require a constant RHS.
+  ConstantSDNode *C = dyn_cast<ConstantSDNode>(N1);
   if (!C)
     return SDValue();
 
-  int64_t MulAmt = C->getSExtValue();
-  unsigned ShiftAmt = llvm::countr_zero<uint64_t>(MulAmt);
+  const APInt &ConstValue = C->getAPIntValue();
 
-  ShiftAmt = ShiftAmt & (32 - 1);
-  SDValue V = N->getOperand(0);
-  SDLoc DL(N);
-
-  SDValue Res;
-  MulAmt >>= ShiftAmt;
-
-  if (MulAmt >= 0) {
-    if (llvm::has_single_bit<uint32_t>(MulAmt - 1)) {
-      // (mul x, 2^N + 1) => (add (shl x, N), x)
-      Res = DAG.getNode(ISD::ADD, DL, VT,
-                        V,
-                        DAG.getNode(ISD::SHL, DL, VT,
-                                    V,
-                                    DAG.getConstant(Log2_32(MulAmt - 1), DL,
-                                                    MVT::i32)));
-    } else if (llvm::has_single_bit<uint32_t>(MulAmt + 1)) {
-      // (mul x, 2^N - 1) => (sub (shl x, N), x)
-      Res = DAG.getNode(ISD::SUB, DL, VT,
-                        DAG.getNode(ISD::SHL, DL, VT,
-                                    V,
-                                    DAG.getConstant(Log2_32(MulAmt + 1), DL,
-                                                    MVT::i32)),
-                        V);
-    } else
+  unsigned TrailingZeroes = ConstValue.countr_zero();
+  if (TrailingZeroes && !Subtarget->isThumb()) {
+    // Conservatively do not lower to shift+add+shift if the mul might be
+    // folded into smul or umul.
+    if (N0->hasOneUse() && (isSignExtended(N0.getNode(), DAG) ||
+                            isZeroExtended(N0.getNode(), DAG)))
       return SDValue();
-  } else {
-    uint64_t MulAmtAbs = -MulAmt;
-    if (llvm::has_single_bit<uint32_t>(MulAmtAbs + 1)) {
-      // (mul x, -(2^N - 1)) => (sub x, (shl x, N))
-      Res = DAG.getNode(ISD::SUB, DL, VT,
-                        V,
-                        DAG.getNode(ISD::SHL, DL, VT,
-                                    V,
-                                    DAG.getConstant(Log2_32(MulAmtAbs + 1), DL,
-                                                    MVT::i32)));
-    } else if (llvm::has_single_bit<uint32_t>(MulAmtAbs - 1)) {
-      // (mul x, -(2^N + 1)) => - (add (shl x, N), x)
-      Res = DAG.getNode(ISD::ADD, DL, VT,
-                        V,
-                        DAG.getNode(ISD::SHL, DL, VT,
-                                    V,
-                                    DAG.getConstant(Log2_32(MulAmtAbs - 1), DL,
-                                                    MVT::i32)));
-      Res = DAG.getNode(ISD::SUB, DL, VT,
-                        DAG.getConstant(0, DL, MVT::i32), Res);
-    } else
+    // Conservatively do not lower to shift+add+shift if the mul might be
+    // folded into madd or msub.
+    if (N->hasOneUse() && (N->user_begin()->getOpcode() == ISD::ADD ||
+                           N->user_begin()->getOpcode() == ISD::SUB))
       return SDValue();
   }
 
-  if (ShiftAmt != 0)
-    Res = DAG.getNode(ISD::SHL, DL, VT,
-                      Res, DAG.getConstant(ShiftAmt, DL, MVT::i32));
+  // Use ShiftedConstValue instead of ConstValue to support both shift+add/sub
+  // and shift+add+shift.
+  APInt ShiftedConstValue = ConstValue.ashr(TrailingZeroes);
+  unsigned ShiftAmt;
+
+  auto Shl = [&](SDValue N0, unsigned N1) {
+    if (!N0.getNode())
+      return SDValue();
+    // If shift causes overflow, ignore this combine.
+    if (N1 >= N0.getValueSizeInBits())
+      return SDValue();
+    SDValue RHS = DAG.getConstant(N1, DL, MVT::i32);
+    return DAG.getNode(ISD::SHL, DL, VT, N0, RHS);
+  };
+  auto Add = [&](SDValue N0, SDValue N1) {
+    if (!N0.getNode() || !N1.getNode())
+      return SDValue();
+    return DAG.getNode(ISD::ADD, DL, VT, N0, N1);
+  };
+  auto Sub = [&](SDValue N0, SDValue N1) {
+    if (!N0.getNode() || !N1.getNode())
+      return SDValue();
+    return DAG.getNode(ISD::SUB, DL, VT, N0, N1);
+  };
+  auto Negate = [&](SDValue N) {
+    if (!N0.getNode())
+      return SDValue();
+    SDValue Zero = DAG.getConstant(0, DL, VT);
+    return DAG.getNode(ISD::SUB, DL, VT, Zero, N);
+  };
+
+  // Can the const C be decomposed into (1+2^M1)*(1+2^N1), eg:
+  // C = 45 is equal to (1+4)*(1+8), we don't decompose it into (1+2)*(16-1) as
+  // the (2^N - 1) can't be execused via a single instruction.
+  auto isPowPlusPlusConst = [](APInt C, APInt &M, APInt &N) {
+    unsigned BitWidth = C.getBitWidth();
+    for (unsigned i = 1; i < BitWidth / 2; i++) {
+      APInt Rem;
+      APInt X(BitWidth, (1 << i) + 1);
+      APInt::sdivrem(C, X, N, Rem);
+      APInt NVMinus1 = N - 1;
+      if (Rem == 0 && NVMinus1.isPowerOf2()) {
+        M = X;
+        return true;
+      }
+    }
+    return false;
+  };
+
+  // Can the const C be decomposed into (2^M + 1) * 2^N + 1), eg:
+  // C = 11 is equal to (1+4)*2+1, we don't decompose it into (1+2)*4-1 as
+  // the (2^N - 1) can't be execused via a single instruction.
+  auto isPowPlusPlusOneConst = [](APInt C, APInt &M, APInt &N) {
+    APInt CVMinus1 = C - 1;
+    if (CVMinus1.isNegative())
+      return false;
+    unsigned TrailingZeroes = CVMinus1.countr_zero();
+    APInt SCVMinus1 = CVMinus1.ashr(TrailingZeroes) - 1;
+    if (SCVMinus1.isPowerOf2()) {
+      unsigned BitWidth = SCVMinus1.getBitWidth();
+      M = APInt(BitWidth, SCVMinus1.logBase2());
+      N = APInt(BitWidth, TrailingZeroes);
+      return true;
+    }
+    return false;
+  };
+
+  // Can the const C be decomposed into (1 - (1 - 2^M) * 2^N), eg:
+  // C = 29 is equal to 1 - (1 - 2^3) * 2^2.
+  auto isPowMinusMinusOneConst = [](APInt C, APInt &M, APInt &N) {
+    APInt CVMinus1 = C - 1;
+    if (CVMinus1.isNegative())
+      return false;
+    unsigned TrailingZeroes = CVMinus1.countr_zero();
+    APInt CVPlus1 = CVMinus1.ashr(TrailingZeroes) + 1;
+    if (CVPlus1.isPowerOf2()) {
+      unsigned BitWidth = CVPlus1.getBitWidth();
+      M = APInt(BitWidth, CVPlus1.logBase2());
+      N = APInt(BitWidth, TrailingZeroes);
+      return true;
+    }
+    return false;
+  };
+
+  if (ConstValue.isNonNegative()) {
+    // (mul x, (2^N + 1) * 2^M) => (shl (add x, (shl x, N)), M)
+    // (mul x, 2^N - 1) => (sub (shl x, N), x)
+    // (mul x, (2^(N-M) - 1) * 2^M) => (sub (shl x, N), (shl x, M))
+    // (mul x, (2^M + 1) * (2^N + 1))
+    //     => MV = (add x, (shl x, M)); (add MV, (shl MV, N))
+    // (mul x, (2^M + 1) * 2^N + 1))
+    //     =>  MV = (add x, (shl x, M)); (add x, (shl MV, N)))
+    // (mul x, 1 - (1 - 2^M) * 2^N))
+    //     =>  MV = sub (x - (shl x, M)); sub (x - (shl MV, N))
+    APInt SCVMinus1 = ShiftedConstValue - 1;
+    APInt SCVPlus1 = ShiftedConstValue + 1;
+    APInt CVPlus1 = ConstValue + 1;
+    APInt CVM, CVN;
+    if (SCVMinus1.isPowerOf2()) {
+      ShiftAmt = SCVMinus1.logBase2();
+      return Shl(Add(N0, Shl(N0, ShiftAmt)), TrailingZeroes);
+    } else if (CVPlus1.isPowerOf2()) {
+      ShiftAmt = CVPlus1.logBase2();
+      return Sub(Shl(N0, ShiftAmt), N0);
+    } else if (SCVPlus1.isPowerOf2()) {
+      ShiftAmt = SCVPlus1.logBase2() + TrailingZeroes;
+      return Sub(Shl(N0, ShiftAmt), Shl(N0, TrailingZeroes));
+    }
+    if (isPowPlusPlusConst(ConstValue, CVM, CVN)) {
+      APInt CVMMinus1 = CVM - 1;
+      APInt CVNMinus1 = CVN - 1;
+      unsigned ShiftM1 = CVMMinus1.logBase2();
+      unsigned ShiftN1 = CVNMinus1.logBase2();
+
+      SDValue MVal = Add(N0, Shl(N0, ShiftM1));
+      return Add(MVal, Shl(MVal, ShiftN1));
+    }
+    if (isPowPlusPlusOneConst(ConstValue, CVM, CVN)) {
+
+      SDValue MVal = Add(N0, Shl(N0, CVM.getZExtValue()));
+      return Add(N0, Shl(MVal, CVN.getZExtValue()));
+    }
+
+    if (isPowMinusMinusOneConst(ConstValue, CVM, CVN)) {
+      SDValue MVal = Sub(N0, Shl(N0, CVM.getZExtValue()));
+      return Sub(N0, Shl(MVal, CVN.getZExtValue()));
+    }
+  } else {
+    // (mul x, -(2^N - 1)) => (sub x, (shl x, N))
+    // (mul x, -(2^N + 1)) => - (add x, (shl x, N)))
+    // (mul x, -(2^(N-M) - 1) * 2^M) => (sub (shl x, M), (shl x, N))
+    APInt SCVPlus1 = -ShiftedConstValue + 1;
+    APInt CVNegPlus1 = -ConstValue + 1;
+    APInt CVNegMinus1 = -ConstValue - 1;
+    if (CVNegPlus1.isPowerOf2()) {
+      ShiftAmt = CVNegPlus1.logBase2();
+      return Sub(N0, Shl(N0, ShiftAmt));
+    } else if (CVNegMinus1.isPowerOf2()) {
+      ShiftAmt = CVNegMinus1.logBase2();
+      return Negate(Add(N0, Shl(N0, ShiftAmt)));
+    } else if (SCVPlus1.isPowerOf2()) {
+      ShiftAmt = SCVPlus1.logBase2() + TrailingZeroes;
+      return Sub(Shl(N0, TrailingZeroes), Shl(N0, ShiftAmt));
+    }
+  }
 
-  // Do not add new nodes to DAG combiner worklist.
-  DCI.CombineTo(N, Res, false);
   return SDValue();
 }
 
@@ -18955,7 +19089,7 @@ SDValue ARMTargetLowering::PerformDAGCombine(SDNode *N,
   case ARMISD::UMLAL:   return PerformUMLALCombine(N, DCI.DAG, Subtarget);
   case ISD::ADD:        return PerformADDCombine(N, DCI, Subtarget);
   case ISD::SUB:        return PerformSUBCombine(N, DCI, Subtarget);
-  case ISD::MUL:        return PerformMULCombine(N, DCI, Subtarget);
+  case ISD::MUL:        return PerformMULCombine(N, DCI.DAG, DCI, Subtarget);
   case ISD::OR:         return PerformORCombine(N, DCI, Subtarget);
   case ISD::XOR:        return PerformXORCombine(N, DCI, Subtarget);
   case ISD::AND:        return PerformANDCombine(N, DCI, Subtarget);
diff --git a/llvm/test/CodeGen/AArch64/load-insert-zero.ll b/llvm/test/CodeGen/AArch64/load-insert-zero.ll
index d6150b6dc45850..cbb2d6081acaed 100644
--- a/llvm/test/CodeGen/AArch64/load-insert-zero.ll
+++ b/llvm/test/CodeGen/AArch64/load-insert-zero.ll
@@ -793,7 +793,7 @@ define void @predictor_4x4_neon(ptr nocapture noundef writeonly %0, i64 noundef
 ; CHECK-NEXT:    ext v2.8b, v2.8b, v0.8b, #1
 ; CHECK-NEXT:    ext v1.8b, v3.8b, v0.8b, #1
 ; CHECK-NEXT:    str s2, [x0, x8]
-; CHECK-NEXT:    add x8, x8, x1
+; CHECK-NEXT:    add x8, x1, x8
 ; CHECK-NEXT:    str s1, [x0, x8]
 ; CHECK-NEXT:    ret
   %5 = load i32, ptr %2, align 4
@@ -856,7 +856,7 @@ define void @predictor_4x4_neon_new(ptr nocapture noundef writeonly %0, i64 noun
 ; CHECK-NEXT:    ldur s3, [x2, #3]
 ; CHECK-NEXT:    uaddl v4.8h, v1.8b, v0.8b
 ; CHECK-NEXT:    urhadd v0.8b, v0.8b, v1.8b
-; CHECK-NEXT:    add x9, x8, x1
+; CHECK-NEXT:    add x9, x1, x8
 ; CHECK-NEXT:    uaddl v5.8h, v2.8b, v1.8b
 ; CHECK-NEXT:    uaddl v3.8h, v3.8b, v2.8b
 ; CHECK-NEXT:    urhadd v1.8b, v1.8b, v2.8b
diff --git a/llvm/test/CodeGen/AArch64/sme2-intrinsics-int-dots.ll b/llvm/test/CodeGen/AArch64/sme2-intrinsics-int-dots.ll
index 111b3fde29a378..33b5a8f5eac530 100644
--- a/llvm/test/CodeGen/AArch64/sme2-intrinsics-int-dots.ll
+++ b/llvm/test/CodeGen/AArch64/sme2-intrinsics-int-dots.ll
@@ -77,7 +77,7 @@ define void @udot_multi_za32_u16_vg1x4_tuple(i64 %stride, ptr %ptr) #1 {
 ; CHECK-NEXT:    mov w8, wzr
 ; CHECK-NEXT:    ld1b { z16.b, z20.b, z24.b, z28.b }, pn8/z, [x1]
 ; CHECK-NEXT:    ld1b { z17.b, z21.b, z25.b, z29.b }, pn8/z, [x1, x0]
-; CHECK-NEXT:    add x10, x9, x0
+; CHECK-NEXT:    add x10, x0, x9
 ; CHECK-NEXT:    ld1b { z18.b, z22.b, z26.b, z30.b }, pn8/z, [x1, x9]
 ; CHECK-NEXT:    ld1b { z19.b, z23.b, z27.b, z31.b }, pn8/z, [x1, x10]
 ; CHECK-NEXT:    udot za.s[w8, 0, vgx4], { z16.b - z19.b }, { z20.b - z23.b }
@@ -270,7 +270,7 @@ define void @usdot_multi_za32_u16_vg1x4_tup...
[truncated]

``````````

</details>


https://github.com/llvm/llvm-project/pull/195334


More information about the llvm-commits mailing list