Skip to content

Commit bde23d6

Browse files
lukel97navaneethshan
authored andcommitted
[SelectionDAG] Expand CTTZ_ELTS[_ZERO_POISON] and handle legalization (#188691)
This is a second attempt at "[SelectionDAG] Expand CTTZ_ELTS[_ZERO_POISON] and handle splitting" (#188220) That PR had to be reverted in 7d39664 because we had crashes on AMDGPU since we didn't have scalarization support, and other crashes on PowerPC because we didn't handle the case when a vector needed widened. Tests for these are added in AMDGPU/cttz-elts.ll, RISCV/rvv/cttz-elts-scalarize.ll and PowerPC/cttz-elts.ll. The former crash has been fixed by adding DAGTypeLegalizer::ScalarizeVecOp_CTTZ_ELTS. The second crash has been fixed by reworking TargetLowering::expandCttzElts. The expansion for CTTZ_ELTS is nearly identical to VECTOR_FIND_LAST_ACTIVE, except it uses a reverse step vector and subtracts the result from VF. The easiest way to fix these crashes without introducing regressions is to reuse the VECTOR_FIND_LAST_ACTIVE expansion which already handles the case where the vector needs widened. This means that the node now needs to take in a boolean vector argument and uses VSELECT instead of an AND to zero out inactive lanes, so the op promotion code has also been shared. (cherry picked from commit 598f353)
1 parent 7e4e25a commit bde23d6

20 files changed

Lines changed: 1175 additions & 566 deletions

‎llvm/include/llvm/CodeGen/BasicTTIImpl.h‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2138,7 +2138,8 @@ class BasicTTIImplBase : public TargetTransformInfoImplCRTPBase<T> {
21382138
VScaleRange = getVScaleRange(I->getCaller(), 64);
21392139

21402140
unsigned EltWidth = getTLI()->getBitWidthForCttzElements(
2141-
RetTy, ArgType.getVectorElementCount(), ZeroIsPoison, &VScaleRange);
2141+
getTLI()->getValueType(DL, RetTy), ArgType.getVectorElementCount(),
2142+
ZeroIsPoison, &VScaleRange);
21422143
Type *NewEltTy = IntegerType::getIntNTy(RetTy->getContext(), EltWidth);
21432144

21442145
// Create the new vector type & get the vector length

‎llvm/include/llvm/CodeGen/ISDOpcodes.h‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1580,7 +1580,7 @@ enum NodeType {
15801580
EXPERIMENTAL_VECTOR_HISTOGRAM,
15811581

15821582
/// Returns the number of number of trailing (least significant) zero elements
1583-
/// in a vector. Has a single i1 vector operand. The result is poison if the
1583+
/// in a vector. Has a single mask vector operand. The result is poison if the
15841584
/// return type isn't wide enough to hold the maximum number of elements in
15851585
/// the input vector.
15861586
CTTZ_ELTS,

‎llvm/include/llvm/CodeGen/TargetLowering.h‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -498,7 +498,7 @@ class LLVM_ABI TargetLoweringBase {
498498

499499
/// Return the minimum number of bits required to hold the maximum possible
500500
/// number of trailing zero vector elements.
501-
unsigned getBitWidthForCttzElements(Type *RetTy, ElementCount EC,
501+
unsigned getBitWidthForCttzElements(EVT RetVT, ElementCount EC,
502502
bool ZeroIsPoison,
503503
const ConstantRange *VScaleRange) const;
504504

@@ -5870,6 +5870,10 @@ class LLVM_ABI TargetLowering : public TargetLoweringBase {
58705870
/// temporarily, advance store position, before re-loading the final vector.
58715871
SDValue expandVECTOR_COMPRESS(SDNode *Node, SelectionDAG &DAG) const;
58725872

5873+
/// Expand a CTTZ_ELTS or CTTZ_ELTS_ZERO_POISON by calculating (VL - i) for
5874+
/// each active lane (i), getting the maximum and subtracting it from VL.
5875+
SDValue expandCttzElts(SDNode *Node, SelectionDAG &DAG) const;
5876+
58735877
/// Expands PARTIAL_REDUCE_S/UMLA nodes to a series of simpler operations,
58745878
/// consisting of zext/sext, extract_subvector, mul and add operations.
58755879
SDValue expandPartialReduceMLA(SDNode *Node, SelectionDAG &DAG) const;

‎llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp‎

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2164,7 +2164,9 @@ bool DAGTypeLegalizer::PromoteIntegerOperand(SDNode *N, unsigned OpNo) {
21642164
Res = PromoteIntOp_VECTOR_HISTOGRAM(N, OpNo);
21652165
break;
21662166
case ISD::VECTOR_FIND_LAST_ACTIVE:
2167-
Res = PromoteIntOp_VECTOR_FIND_LAST_ACTIVE(N, OpNo);
2167+
case ISD::CTTZ_ELTS:
2168+
case ISD::CTTZ_ELTS_ZERO_POISON:
2169+
Res = PromoteIntOp_UnaryBooleanVectorOp(N, OpNo);
21682170
break;
21692171
case ISD::GET_ACTIVE_LANE_MASK:
21702172
Res = PromoteIntOp_GET_ACTIVE_LANE_MASK(N);
@@ -2992,8 +2994,8 @@ SDValue DAGTypeLegalizer::PromoteIntOp_VECTOR_HISTOGRAM(SDNode *N,
29922994
return SDValue(DAG.UpdateNodeOperands(N, NewOps), 0);
29932995
}
29942996

2995-
SDValue DAGTypeLegalizer::PromoteIntOp_VECTOR_FIND_LAST_ACTIVE(SDNode *N,
2996-
unsigned OpNo) {
2997+
SDValue DAGTypeLegalizer::PromoteIntOp_UnaryBooleanVectorOp(SDNode *N,
2998+
unsigned OpNo) {
29972999
assert(OpNo == 0 && "Unexpected operand for promotion");
29983000
SDValue Op = N->getOperand(0);
29993001

‎llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -418,7 +418,7 @@ class LLVM_LIBRARY_VISIBILITY DAGTypeLegalizer {
418418
SDValue PromoteIntOp_VP_STRIDED(SDNode *N, unsigned OpNo);
419419
SDValue PromoteIntOp_VP_SPLICE(SDNode *N, unsigned OpNo);
420420
SDValue PromoteIntOp_VECTOR_HISTOGRAM(SDNode *N, unsigned OpNo);
421-
SDValue PromoteIntOp_VECTOR_FIND_LAST_ACTIVE(SDNode *N, unsigned OpNo);
421+
SDValue PromoteIntOp_UnaryBooleanVectorOp(SDNode *N, unsigned OpNo);
422422
SDValue PromoteIntOp_GET_ACTIVE_LANE_MASK(SDNode *N);
423423
SDValue PromoteIntOp_PARTIAL_REDUCE_MLA(SDNode *N);
424424
SDValue PromoteIntOp_LOOP_DEPENDENCE_MASK(SDNode *N, unsigned OpNo);
@@ -881,6 +881,7 @@ class LLVM_LIBRARY_VISIBILITY DAGTypeLegalizer {
881881
SDValue ScalarizeVecOp_CMP(SDNode *N);
882882
SDValue ScalarizeVecOp_FAKE_USE(SDNode *N);
883883
SDValue ScalarizeVecOp_VECTOR_FIND_LAST_ACTIVE(SDNode *N);
884+
SDValue ScalarizeVecOp_CTTZ_ELTS(SDNode *N);
884885

885886
//===--------------------------------------------------------------------===//
886887
// Vector Splitting Support: LegalizeVectorTypes.cpp
@@ -987,6 +988,7 @@ class LLVM_LIBRARY_VISIBILITY DAGTypeLegalizer {
987988
SDValue SplitVecOp_FPOpDifferentTypes(SDNode *N);
988989
SDValue SplitVecOp_CMP(SDNode *N);
989990
SDValue SplitVecOp_FP_TO_XINT_SAT(SDNode *N);
991+
SDValue SplitVecOp_CttzElts(SDNode *N);
990992
SDValue SplitVecOp_VP_CttzElements(SDNode *N);
991993
SDValue SplitVecOp_VECTOR_HISTOGRAM(SDNode *N);
992994
SDValue SplitVecOp_PARTIAL_REDUCE_MLA(SDNode *N);

‎llvm/lib/CodeGen/SelectionDAG/LegalizeVectorOps.cpp‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -518,6 +518,8 @@ SDValue VectorLegalizer::LegalizeOp(SDValue Op) {
518518
case ISD::VECREDUCE_FMIN:
519519
case ISD::VECREDUCE_FMINIMUM:
520520
case ISD::VECREDUCE_FMUL:
521+
case ISD::CTTZ_ELTS:
522+
case ISD::CTTZ_ELTS_ZERO_POISON:
521523
case ISD::VECTOR_FIND_LAST_ACTIVE:
522524
Action = TLI.getOperationAction(Node->getOpcode(),
523525
Node->getOperand(0).getValueType());
@@ -1355,6 +1357,10 @@ void VectorLegalizer::Expand(SDNode *Node, SmallVectorImpl<SDValue> &Results) {
13551357
case ISD::VECTOR_COMPRESS:
13561358
Results.push_back(TLI.expandVECTOR_COMPRESS(Node, DAG));
13571359
return;
1360+
case ISD::CTTZ_ELTS:
1361+
case ISD::CTTZ_ELTS_ZERO_POISON:
1362+
Results.push_back(TLI.expandCttzElts(Node, DAG));
1363+
return;
13581364
case ISD::VECTOR_FIND_LAST_ACTIVE:
13591365
Results.push_back(TLI.expandVectorFindLastActive(Node, DAG));
13601366
return;

‎llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp‎

Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -909,6 +909,10 @@ bool DAGTypeLegalizer::ScalarizeVectorOperand(SDNode *N, unsigned OpNo) {
909909
case ISD::VECTOR_FIND_LAST_ACTIVE:
910910
Res = ScalarizeVecOp_VECTOR_FIND_LAST_ACTIVE(N);
911911
break;
912+
case ISD::CTTZ_ELTS:
913+
case ISD::CTTZ_ELTS_ZERO_POISON:
914+
Res = ScalarizeVecOp_CTTZ_ELTS(N);
915+
break;
912916
}
913917

914918
// If the result is null, the sub-method took care of registering results etc.
@@ -1221,6 +1225,18 @@ SDValue DAGTypeLegalizer::ScalarizeVecOp_VECTOR_FIND_LAST_ACTIVE(SDNode *N) {
12211225
return DAG.getConstant(0, SDLoc(N), VT);
12221226
}
12231227

1228+
SDValue DAGTypeLegalizer::ScalarizeVecOp_CTTZ_ELTS(SDNode *N) {
1229+
// The number of trailing zero elements is 1 if the element is 0, and 0
1230+
// otherwise.
1231+
if (N->getOpcode() == ISD::CTTZ_ELTS_ZERO_POISON)
1232+
return DAG.getConstant(0, SDLoc(N), N->getValueType(0));
1233+
SDValue Op = GetScalarizedVector(N->getOperand(0));
1234+
SDValue SetCC =
1235+
DAG.getSetCC(SDLoc(N), MVT::i1, Op,
1236+
DAG.getConstant(0, SDLoc(N), Op.getValueType()), ISD::SETEQ);
1237+
return DAG.getZExtOrTrunc(SetCC, SDLoc(N), N->getValueType(0));
1238+
}
1239+
12241240
//===----------------------------------------------------------------------===//
12251241
// Result Vector Splitting
12261242
//===----------------------------------------------------------------------===//
@@ -3742,6 +3758,10 @@ bool DAGTypeLegalizer::SplitVectorOperand(SDNode *N, unsigned OpNo) {
37423758
case ISD::VP_REDUCE_FMINIMUM:
37433759
Res = SplitVecOp_VP_REDUCE(N, OpNo);
37443760
break;
3761+
case ISD::CTTZ_ELTS:
3762+
case ISD::CTTZ_ELTS_ZERO_POISON:
3763+
Res = SplitVecOp_CttzElts(N);
3764+
break;
37453765
case ISD::VP_CTTZ_ELTS:
37463766
case ISD::VP_CTTZ_ELTS_ZERO_UNDEF:
37473767
Res = SplitVecOp_VP_CttzElements(N);
@@ -4828,6 +4848,26 @@ SDValue DAGTypeLegalizer::SplitVecOp_FP_TO_XINT_SAT(SDNode *N) {
48284848
return DAG.getNode(ISD::CONCAT_VECTORS, dl, ResVT, Lo, Hi);
48294849
}
48304850

4851+
SDValue DAGTypeLegalizer::SplitVecOp_CttzElts(SDNode *N) {
4852+
SDLoc DL(N);
4853+
EVT ResVT = N->getValueType(0);
4854+
4855+
SDValue Lo, Hi;
4856+
SDValue VecOp = N->getOperand(0);
4857+
GetSplitVector(VecOp, Lo, Hi);
4858+
4859+
// if CTTZ_ELTS(Lo) != VL => CTTZ_ELTS(Lo).
4860+
// else => VL + (CTTZ_ELTS(Hi) or CTTZ_ELTS_ZERO_POISON(Hi)).
4861+
SDValue ResLo = DAG.getNode(ISD::CTTZ_ELTS, DL, ResVT, Lo);
4862+
SDValue VL =
4863+
DAG.getElementCount(DL, ResVT, Lo.getValueType().getVectorElementCount());
4864+
SDValue ResLoNotVL =
4865+
DAG.getSetCC(DL, getSetCCResultType(ResVT), ResLo, VL, ISD::SETNE);
4866+
SDValue ResHi = DAG.getNode(N->getOpcode(), DL, ResVT, Hi);
4867+
return DAG.getSelect(DL, ResVT, ResLoNotVL, ResLo,
4868+
DAG.getNode(ISD::ADD, DL, ResVT, VL, ResHi));
4869+
}
4870+
48314871
SDValue DAGTypeLegalizer::SplitVecOp_VP_CttzElements(SDNode *N) {
48324872
SDLoc DL(N);
48334873
EVT ResVT = N->getValueType(0);

‎llvm/lib/CodeGen/SelectionDAG/SelectionDAGBuilder.cpp‎

Lines changed: 10 additions & 46 deletions
Original file line numberDiff line numberDiff line change
@@ -8324,57 +8324,21 @@ void SelectionDAGBuilder::visitIntrinsicCall(const CallInst &I,
83248324
return;
83258325
}
83268326
case Intrinsic::experimental_cttz_elts: {
8327-
auto DL = getCurSDLoc();
83288327
SDValue Op = getValue(I.getOperand(0));
83298328
EVT OpVT = Op.getValueType();
83308329
EVT RetTy = TLI.getValueType(DAG.getDataLayout(), I.getType());
83318330
bool ZeroIsPoison =
83328331
!cast<ConstantSDNode>(getValue(I.getOperand(1)))->isZero();
8333-
8334-
if (!TLI.shouldExpandCttzElements(OpVT)) {
8335-
SDValue Ret = DAG.getNode(ZeroIsPoison ? ISD::CTTZ_ELTS_ZERO_POISON
8336-
: ISD::CTTZ_ELTS,
8337-
sdl, RetTy, Op);
8338-
setValue(&I, Ret);
8339-
return;
8340-
}
8341-
8342-
if (OpVT.getScalarType() != MVT::i1) {
8343-
// Compare the input vector elements to zero & use to count trailing zeros
8344-
SDValue AllZero = DAG.getConstant(0, DL, OpVT);
8345-
OpVT = EVT::getVectorVT(*DAG.getContext(), MVT::i1,
8346-
OpVT.getVectorElementCount());
8347-
Op = DAG.getSetCC(DL, OpVT, Op, AllZero, ISD::SETNE);
8348-
}
8349-
8350-
// If the zero-is-poison flag is set, we can assume the upper limit
8351-
// of the result is VF-1.
8352-
ConstantRange VScaleRange(1, true); // Dummy value.
8353-
if (isa<ScalableVectorType>(I.getOperand(0)->getType()))
8354-
VScaleRange = getVScaleRange(I.getCaller(), 64);
8355-
unsigned EltWidth = TLI.getBitWidthForCttzElements(
8356-
I.getType(), OpVT.getVectorElementCount(), ZeroIsPoison, &VScaleRange);
8357-
8358-
MVT NewEltTy = MVT::getIntegerVT(EltWidth);
8359-
8360-
// Create the new vector type & get the vector length
8361-
EVT NewVT = EVT::getVectorVT(*DAG.getContext(), NewEltTy,
8362-
OpVT.getVectorElementCount());
8363-
8364-
SDValue VL =
8365-
DAG.getElementCount(DL, NewEltTy, OpVT.getVectorElementCount());
8366-
8367-
SDValue StepVec = DAG.getStepVector(DL, NewVT);
8368-
SDValue SplatVL = DAG.getSplat(NewVT, DL, VL);
8369-
SDValue StepVL = DAG.getNode(ISD::SUB, DL, NewVT, SplatVL, StepVec);
8370-
SDValue Ext = DAG.getNode(ISD::SIGN_EXTEND, DL, NewVT, Op);
8371-
SDValue And = DAG.getNode(ISD::AND, DL, NewVT, StepVL, Ext);
8372-
SDValue Max = DAG.getNode(ISD::VECREDUCE_UMAX, DL, NewEltTy, And);
8373-
SDValue Sub = DAG.getNode(ISD::SUB, DL, NewEltTy, VL, Max);
8374-
8375-
SDValue Ret = DAG.getZExtOrTrunc(Sub, DL, RetTy);
8376-
8377-
setValue(&I, Ret);
8332+
if (OpVT.getVectorElementType() != MVT::i1) {
8333+
// Compare the input vector elements to zero & use to count trailing
8334+
// zeros.
8335+
SDValue AllZero = DAG.getConstant(0, sdl, OpVT);
8336+
EVT I1OpVT = OpVT.changeVectorElementType(*DAG.getContext(), MVT::i1);
8337+
Op = DAG.getSetCC(sdl, I1OpVT, Op, AllZero, ISD::SETNE);
8338+
}
8339+
setValue(&I, DAG.getNode(ZeroIsPoison ? ISD::CTTZ_ELTS_ZERO_POISON
8340+
: ISD::CTTZ_ELTS,
8341+
sdl, RetTy, Op));
83788342
return;
83798343
}
83808344
case Intrinsic::vector_insert: {

‎llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp‎

Lines changed: 48 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -10120,22 +10120,24 @@ SDValue TargetLowering::expandVPCTTZElements(SDNode *N,
1012010120
return DAG.getNode(ISD::VP_REDUCE_UMIN, DL, ResVT, ExtEVL, Select, Mask, EVL);
1012110121
}
1012210122

10123-
SDValue TargetLowering::expandVectorFindLastActive(SDNode *N,
10124-
SelectionDAG &DAG) const {
10125-
SDLoc DL(N);
10126-
SDValue Mask = N->getOperand(0);
10123+
/// Returns a type-legalized version of \p Mask as the first item in the
10124+
/// pair. The second item contains a type-legalized step vector that's
10125+
/// guaranteed to fit the number of elements in \p Mask.
10126+
static std::pair<SDValue, SDValue>
10127+
getLegalMaskAndStepVector(SDValue Mask, bool ZeroIsPoison, SDLoc DL,
10128+
SelectionDAG &DAG) {
1012710129
EVT MaskVT = Mask.getValueType();
1012810130
EVT BoolVT = MaskVT.getScalarType();
1012910131

1013010132
// Find a suitable type for a stepvector.
10133+
// If zero is poison, we can assume the upper limit of the result is VF-1.
1013110134
ConstantRange VScaleRange(1, /*isFullSet=*/true); // Fixed length default.
1013210135
if (MaskVT.isScalableVector())
1013310136
VScaleRange = getVScaleRange(&DAG.getMachineFunction().getFunction(), 64);
1013410137
const TargetLowering &TLI = DAG.getTargetLoweringInfo();
1013510138
uint64_t EltWidth = TLI.getBitWidthForCttzElements(
10136-
EVT(getVectorIdxTy(DAG.getDataLayout())).getTypeForEVT(*DAG.getContext()),
10137-
MaskVT.getVectorElementCount(),
10138-
/*ZeroIsPoison=*/true, &VScaleRange);
10139+
EVT(TLI.getVectorIdxTy(DAG.getDataLayout())),
10140+
MaskVT.getVectorElementCount(), ZeroIsPoison, &VScaleRange);
1013910141
// If the step vector element type is smaller than the mask element type,
1014010142
// use the mask type directly to avoid widening issues.
1014110143
EltWidth = std::max(EltWidth, BoolVT.getFixedSizeInBits());
@@ -10151,7 +10153,6 @@ SDValue TargetLowering::expandVectorFindLastActive(SDNode *N,
1015110153
SDValue StepVec;
1015210154
if (TypeAction == TargetLowering::TypePromoteInteger) {
1015310155
StepVecVT = TLI.getTypeToTransformTo(*DAG.getContext(), StepVecVT);
10154-
StepVT = StepVecVT.getVectorElementType();
1015510156
StepVec = DAG.getStepVector(DL, StepVecVT);
1015610157
} else if (TypeAction == TargetLowering::TypeWidenVector) {
1015710158
// For widening, the element count changes. Create a step vector with only
@@ -10168,13 +10169,21 @@ SDValue TargetLowering::expandVectorFindLastActive(SDNode *N,
1016810169
EVT WideMaskVT = EVT::getVectorVT(*DAG.getContext(), BoolVT, WideNumElts);
1016910170
SDValue ZeroMask = DAG.getConstant(0, DL, WideMaskVT);
1017010171
Mask = DAG.getInsertSubvector(DL, ZeroMask, Mask, 0);
10171-
10172-
StepVecVT = WideVecVT;
10173-
StepVT = WideVecVT.getVectorElementType();
1017410172
} else {
1017510173
StepVec = DAG.getStepVector(DL, StepVecVT);
1017610174
}
1017710175

10176+
return {Mask, StepVec};
10177+
}
10178+
10179+
SDValue TargetLowering::expandVectorFindLastActive(SDNode *N,
10180+
SelectionDAG &DAG) const {
10181+
SDLoc DL(N);
10182+
auto [Mask, StepVec] = getLegalMaskAndStepVector(
10183+
N->getOperand(0), /*ZeroIsPoison=*/true, DL, DAG);
10184+
EVT StepVecVT = StepVec.getValueType();
10185+
EVT StepVT = StepVec.getValueType().getVectorElementType();
10186+
1017810187
// Zero out lanes with inactive elements, then find the highest remaining
1017910188
// value from the stepvector.
1018010189
SDValue Zeroes = DAG.getConstant(0, DL, StepVecVT);
@@ -12596,6 +12605,34 @@ SDValue TargetLowering::expandVECTOR_COMPRESS(SDNode *Node,
1259612605
return DAG.getLoad(VecVT, DL, Chain, StackPtr, PtrInfo);
1259712606
}
1259812607

12608+
SDValue TargetLowering::expandCttzElts(SDNode *Node, SelectionDAG &DAG) const {
12609+
SDLoc DL(Node);
12610+
EVT VT = Node->getValueType(0);
12611+
12612+
bool ZeroIsPoison = Node->getOpcode() == ISD::CTTZ_ELTS_ZERO_POISON;
12613+
auto [Mask, StepVec] =
12614+
getLegalMaskAndStepVector(Node->getOperand(0), ZeroIsPoison, DL, DAG);
12615+
EVT StepVecVT = StepVec.getValueType();
12616+
EVT StepVT = StepVecVT.getVectorElementType();
12617+
12618+
// Promote the scalar result type early to avoid redundant zexts.
12619+
if (getTypeAction(StepVT.getSimpleVT()) == TypePromoteInteger)
12620+
StepVT = getTypeToTransformTo(*DAG.getContext(), StepVT);
12621+
12622+
SDValue VL =
12623+
DAG.getElementCount(DL, StepVT, StepVecVT.getVectorElementCount());
12624+
SDValue SplatVL = DAG.getSplat(StepVecVT, DL, VL);
12625+
StepVec = DAG.getNode(ISD::SUB, DL, StepVecVT, SplatVL, StepVec);
12626+
SDValue Zeroes = DAG.getConstant(0, DL, StepVecVT);
12627+
SDValue Select = DAG.getSelect(DL, StepVecVT, Mask, StepVec, Zeroes);
12628+
SDValue Max = DAG.getNode(ISD::VECREDUCE_UMAX, DL,
12629+
StepVecVT.getVectorElementType(), Select);
12630+
SDValue Sub = DAG.getNode(ISD::SUB, DL, StepVT, VL,
12631+
DAG.getZExtOrTrunc(Max, DL, StepVT));
12632+
12633+
return DAG.getZExtOrTrunc(Sub, DL, VT);
12634+
}
12635+
1259912636
SDValue TargetLowering::expandPartialReduceMLA(SDNode *N,
1260012637
SelectionDAG &DAG) const {
1260112638
SDLoc DL(N);

‎llvm/lib/CodeGen/TargetLoweringBase.cpp‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1347,7 +1347,7 @@ bool TargetLoweringBase::isFreeAddrSpaceCast(unsigned SrcAS,
13471347
}
13481348

13491349
unsigned TargetLoweringBase::getBitWidthForCttzElements(
1350-
Type *RetTy, ElementCount EC, bool ZeroIsPoison,
1350+
EVT RetVT, ElementCount EC, bool ZeroIsPoison,
13511351
const ConstantRange *VScaleRange) const {
13521352
// Find the smallest "sensible" element type to use for the expansion.
13531353
ConstantRange CR(APInt(64, EC.getKnownMinValue()));
@@ -1357,7 +1357,7 @@ unsigned TargetLoweringBase::getBitWidthForCttzElements(
13571357
if (ZeroIsPoison)
13581358
CR = CR.subtract(APInt(64, 1));
13591359

1360-
unsigned EltWidth = RetTy->getScalarSizeInBits();
1360+
unsigned EltWidth = RetVT.getScalarSizeInBits();
13611361
EltWidth = std::min(EltWidth, CR.getActiveBits());
13621362
EltWidth = std::max(llvm::bit_ceil(EltWidth), (unsigned)8);
13631363

0 commit comments

Comments
 (0)