[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.
This commit is contained in:
parent
7d7cd745af
commit
c1df6937ba
@ -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<unsigned>(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));
|
||||
|
||||
@ -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
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user