From c1df6937bac1191663c7f3a84f5c8faf2e79db35 Mon Sep 17 00:00:00 2001 From: Craig Topper Date: Sat, 21 Mar 2026 09:31:20 -0700 Subject: [PATCH] [TargetLowering] Use legally typed shifts to split chunks in expandDIVREMByConstant. (#187567) This replaces LegalVT with HiLoVT and LegalWidth with HBitWidth as they are the same for all current uses. Then we rewrite the shifts to operate on LL and LH. There's a slight regression on RISC-V due to different node creation order leading to different DAG combine order. I have other refactoring I'd like to explore then I may try to fix that. --- .../CodeGen/SelectionDAG/TargetLowering.cpp | 84 +++++++++++-------- llvm/test/CodeGen/RISCV/urem-lkk.ll | 24 +++--- 2 files changed, 61 insertions(+), 47 deletions(-) diff --git a/llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp b/llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp index c9bddf636c75..4963003233cd 100644 --- a/llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp +++ b/llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp @@ -8246,23 +8246,18 @@ bool TargetLowering::expandDIVREMByConstant(SDNode *N, unsigned BitWidth = VT.getScalarSizeInBits(); unsigned BestChunkWidth = 0; - // Determine the legal scalar integer type for chunk operations. - EVT LegalVT = getTypeToTransformTo(*DAG.getContext(), VT); - unsigned LegalWidth = LegalVT.getScalarSizeInBits(); - unsigned MaxChunk = std::min(LegalWidth, BitWidth); - // Search for I where 2^I % Divisor == 1 - for (unsigned I = MaxChunk, E = MaxChunk / 2; I > E; --I) { + for (unsigned I = HBitWidth - 1, E = HBitWidth / 2; I > E; --I) { APInt Mod = APInt::getOneBitSet(Divisor.getBitWidth(), I).urem(Divisor); if (Mod.isOne()) { - // Ensure (NumChunks * MaxChunkValue) doesn't overflow LegalVT + // Ensure (NumChunks * MaxChunkValue) doesn't overflow HiLoVT unsigned NumChunks = divideCeil(BitWidth, I); - // Ensure the sum won't overflow the hardware register (LegalWidth). + // Ensure the sum won't overflow the hardware register (HBitWidth). // Summing N chunks adds ceil(log2(N)) extra carry bits to the width. // Safety check: Base Chunk Width (I) + Carry Bits <= Register Width. - if (I + llvm::bit_width(NumChunks - 1) <= LegalWidth) { + if (I + llvm::bit_width(NumChunks - 1) <= HBitWidth) { BestChunkWidth = I; break; } @@ -8272,46 +8267,64 @@ bool TargetLowering::expandDIVREMByConstant(SDNode *N, if (!BestChunkWidth) return false; - SDValue In = - LL ? DAG.getNode(ISD::BUILD_PAIR, dl, VT, LL, LH) : N->getOperand(0); + assert(!LL == !LH && "Expected both input halves or no input halves!"); + if (!LL) + std::tie(LL, LH) = DAG.SplitScalar(N->getOperand(0), dl, HiLoVT, HiLoVT); + if (TrailingZeros) { // Save the shifted off bits if we need the remainder. if (Opcode != ISD::UDIV) { - APInt Mask = APInt::getLowBitsSet(BitWidth, TrailingZeros); - PartialRem = - DAG.getNode(ISD::AND, dl, VT, In, DAG.getConstant(Mask, dl, VT)); + APInt Mask = APInt::getLowBitsSet(HBitWidth, TrailingZeros); + PartialRem = DAG.getNode(ISD::AND, dl, HiLoVT, LL, + DAG.getConstant(Mask, dl, HiLoVT)); } - EVT ShiftVT = getShiftAmountTy(VT, DAG.getDataLayout()); - In = DAG.getNode(ISD::SRL, dl, VT, In, - DAG.getShiftAmountConstant(TrailingZeros, ShiftVT, dl)); - std::tie(LL, LH) = DAG.SplitScalar(In, dl, HiLoVT, HiLoVT); - } else if (!LL) { - std::tie(LL, LH) = DAG.SplitScalar(In, dl, HiLoVT, HiLoVT); + if (isOperationLegal(ISD::FSHR, HiLoVT)) + LL = DAG.getNode(ISD::FSHR, dl, HiLoVT, LH, LL, + DAG.getShiftAmountConstant(TrailingZeros, HiLoVT, dl)); + else + LL = DAG.getNode( + ISD::OR, dl, HiLoVT, + DAG.getNode(ISD::SRL, dl, HiLoVT, LL, + DAG.getShiftAmountConstant(TrailingZeros, HiLoVT, dl)), + DAG.getNode(ISD::SHL, dl, HiLoVT, LH, + DAG.getShiftAmountConstant(HBitWidth - TrailingZeros, + HiLoVT, dl))); + LH = DAG.getNode(ISD::SRL, dl, HiLoVT, LH, + DAG.getShiftAmountConstant(TrailingZeros, HiLoVT, dl)); } - SDValue TotalSum = DAG.getConstant(0, dl, LegalVT); + Sum = DAG.getConstant(0, dl, HiLoVT); SDValue Mask = DAG.getConstant( - APInt::getLowBitsSet(LegalWidth, BestChunkWidth), dl, LegalVT); + APInt::getLowBitsSet(HBitWidth, BestChunkWidth), dl, HiLoVT); for (unsigned I = 0; I < BitWidth; I += BestChunkWidth) { - SDValue Shift = DAG.getShiftAmountConstant(I, VT, dl); - SDValue Chunk = DAG.getNode(ISD::SRL, dl, VT, In, Shift); - // Truncate to LegalVT - SDValue TruncChunk = DAG.getNode(ISD::TRUNCATE, dl, LegalVT, Chunk); + SDValue Chunk; + if (I == 0) { + Chunk = LL; + } else if (I >= HBitWidth) { + Chunk = + DAG.getNode(ISD::SRL, dl, HiLoVT, LH, + DAG.getShiftAmountConstant(I - HBitWidth, HiLoVT, dl)); + } else if (isOperationLegal(ISD::FSHR, HiLoVT)) { + Chunk = DAG.getNode(ISD::FSHR, dl, HiLoVT, LH, LL, + DAG.getShiftAmountConstant(I, HiLoVT, dl)); + } else { + Chunk = DAG.getNode( + ISD::OR, dl, HiLoVT, + DAG.getNode(ISD::SRL, dl, HiLoVT, LL, + DAG.getShiftAmountConstant(I, HiLoVT, dl)), + DAG.getNode(ISD::SHL, dl, HiLoVT, LH, + DAG.getShiftAmountConstant(HBitWidth - I, HiLoVT, dl))); + } + // For the last chunk, we might not need a mask if it's smaller than // BestChunkWidth, but applying it is always safe. - SDValue MaskedChunk = - DAG.getNode(ISD::AND, dl, LegalVT, TruncChunk, Mask); - TotalSum = DAG.getNode(ISD::ADD, dl, LegalVT, TotalSum, MaskedChunk); + SDValue MaskedChunk = DAG.getNode(ISD::AND, dl, HiLoVT, Chunk, Mask); + Sum = DAG.getNode(ISD::ADD, dl, HiLoVT, Sum, MaskedChunk); } - Sum = DAG.getNode(ISD::ZERO_EXTEND, dl, HiLoVT, TotalSum); } - // If we didn't find a sum, we can't do the expansion. - if (!Sum) - return false; - // Perform a HiLoVT urem on the Sum using truncated divisor. SDValue RemL = DAG.getNode(ISD::UREM, dl, HiLoVT, Sum, @@ -8346,8 +8359,7 @@ bool TargetLowering::expandDIVREMByConstant(SDNode *N, RemL = DAG.getNode(ISD::SHL, dl, HiLoVT, RemL, DAG.getShiftAmountConstant(TrailingZeros, HiLoVT, dl)); - SDValue PartialRemLo = DAG.getNode(ISD::TRUNCATE, dl, HiLoVT, PartialRem); - RemL = DAG.getNode(ISD::ADD, dl, HiLoVT, RemL, PartialRemLo); + RemL = DAG.getNode(ISD::ADD, dl, HiLoVT, RemL, PartialRem); } Result.push_back(RemL); Result.push_back(DAG.getConstant(0, dl, HiLoVT)); diff --git a/llvm/test/CodeGen/RISCV/urem-lkk.ll b/llvm/test/CodeGen/RISCV/urem-lkk.ll index 00fe1a5eb111..52ef3a6e72ff 100644 --- a/llvm/test/CodeGen/RISCV/urem-lkk.ll +++ b/llvm/test/CodeGen/RISCV/urem-lkk.ll @@ -231,21 +231,23 @@ define i64 @dont_fold_urem_i64(i64 %x) nounwind { ; RV32IM: # %bb.0: ; RV32IM-NEXT: slli a2, a1, 31 ; RV32IM-NEXT: srli a3, a0, 1 -; RV32IM-NEXT: andi a4, a1, 2046 +; RV32IM-NEXT: lui a4, 512 +; RV32IM-NEXT: srli a5, a1, 1 ; RV32IM-NEXT: srli a1, a1, 11 ; RV32IM-NEXT: or a2, a3, a2 -; RV32IM-NEXT: slli a4, a4, 10 +; RV32IM-NEXT: slli a5, a5, 11 ; RV32IM-NEXT: srli a3, a2, 21 -; RV32IM-NEXT: or a3, a3, a4 -; RV32IM-NEXT: lui a4, 21400 -; RV32IM-NEXT: slli a2, a2, 11 -; RV32IM-NEXT: srli a2, a2, 11 -; RV32IM-NEXT: add a2, a2, a3 -; RV32IM-NEXT: li a3, 49 -; RV32IM-NEXT: addi a4, a4, -2006 +; RV32IM-NEXT: or a3, a3, a5 +; RV32IM-NEXT: lui a5, 21400 +; RV32IM-NEXT: addi a4, a4, -1 +; RV32IM-NEXT: and a2, a2, a4 ; RV32IM-NEXT: add a1, a2, a1 -; RV32IM-NEXT: mulhu a2, a1, a4 -; RV32IM-NEXT: mul a2, a2, a3 +; RV32IM-NEXT: li a2, 49 +; RV32IM-NEXT: addi a5, a5, -2006 +; RV32IM-NEXT: and a3, a3, a4 +; RV32IM-NEXT: add a1, a1, a3 +; RV32IM-NEXT: mulhu a3, a1, a5 +; RV32IM-NEXT: mul a2, a3, a2 ; RV32IM-NEXT: sub a1, a1, a2 ; RV32IM-NEXT: slli a1, a1, 1 ; RV32IM-NEXT: andi a0, a0, 1