[llvm] [RISCV] Fold vp.reverse of vp.load through binary ops (PR #205529)
Min-Yih Hsu via llvm-commits
llvm-commits at lists.llvm.org
Thu Jun 25 09:20:49 PDT 2026
================
@@ -19649,61 +19657,84 @@ static SDValue performReverseEVLCombine(SDNode *N, SelectionDAG &DAG,
//
// splice.right(reverse(vp.load(ADDR, REVMASK, EVL)), poison, EVL)
// -> vp.strided.load(ADDR, -1, MASK, EVL)
-
- // Check if its first operand is a vp.load.
+ //
+ // vp.reverse(binop(vp.load(ADDR, REVMASK, EVL), splat), EVL)
+ // -> binop(vp.strided.load(ADDR, -1, MASK, EVL), splat)
using namespace SDPatternMatch;
SDValue Op, EVL;
- if (!sd_match(N,
- m_ReverseEVL(m_OneUse(m_Value(Op, m_SpecificOpc(ISD::VP_LOAD))),
- m_Value(EVL))))
+ if (!sd_match(N, m_ReverseEVL(m_Value(Op), m_Value(EVL))))
return SDValue();
- auto *VPLoad = cast<VPLoadSDNode>(Op);
-
- EVT LoadVT = VPLoad->getValueType(0);
- // We do not have a strided_load version for masks, and the evl of vp.reverse
- // and vp.load should always be the same.
- if (!LoadVT.getVectorElementType().isByteSized() ||
- EVL != VPLoad->getVectorLength())
- return SDValue();
+ // Check that all leaves are splats or vp_loads, and collect the latter.
+ SmallVector<SDValue> Worklist = {Op};
+ SmallVector<VPLoadSDNode *> VPLoads;
+ while (!Worklist.empty()) {
+ SDValue X = Worklist.pop_back_val();
+ if (!hasOneUser(X))
+ return SDValue();
+ if (auto *VPLoad = dyn_cast<VPLoadSDNode>(X))
+ VPLoads.push_back(VPLoad);
+ else if (DAG.isSplatValue(X))
+ continue;
+ else if (DAG.getTargetLoweringInfo().isBinOp(X.getOpcode()) &&
+ X->getNumValues() == 1)
+ append_range(Worklist, X->op_values());
+ else
+ return SDValue();
+ }
- SDValue LoadMask = VPLoad->getMask();
- // If Mask is all ones, then load is unmasked and can be reversed.
- if (!isOneOrOneSplat(LoadMask)) {
- // If the mask is not all ones, we can reverse the load if the mask was also
- // reversed by a vp.reverse with the same EVL.
- SDValue OrigMask;
- if (!sd_match(LoadMask, m_ReverseEVL(m_Value(OrigMask), m_Specific(EVL))))
+ SmallVector<SDValue> LoadMasks;
+ for (auto *VPLoad : VPLoads) {
+ EVT LoadVT = VPLoad->getValueType(0);
+ // We do not have a strided_load version for masks, and the evl of
+ // vp.reverse and vp.load should always be the same.
+ if (!LoadVT.getVectorElementType().isByteSized() ||
+ EVL != VPLoad->getVectorLength())
return SDValue();
- LoadMask = OrigMask;
+
+ SDValue LoadMask = VPLoad->getMask();
+ // If Mask is all ones, then load is unmasked and can be reversed.
+ if (isOneOrOneSplat(LoadMask)) {
+ LoadMasks.push_back(LoadMask);
+ } else {
+ // If the mask is not all ones, we can reverse the load if the mask was
+ // also reversed by a vp.reverse with the same EVL.
+ SDValue OrigMask;
+ if (!sd_match(LoadMask, m_ReverseEVL(m_Value(OrigMask), m_Specific(EVL))))
+ return SDValue();
+ LoadMasks.push_back(OrigMask);
+ }
}
- // Base = LoadAddr + (NumElem - 1) * ElemWidthByte
+ // Reverse the vp_loads.
SDLoc DL(N);
MVT XLenVT = Subtarget.getXLenVT();
- SDValue NumElem = VPLoad->getVectorLength();
- uint64_t ElemWidthByte = VPLoad->getValueType(0).getScalarSizeInBits() / 8;
-
- SDValue Temp1 = DAG.getNode(ISD::SUB, DL, XLenVT, NumElem,
- DAG.getConstant(1, DL, XLenVT));
+ uint64_t ElemWidthByte = N->getValueType(0).getScalarSizeInBits() / 8;
+ SDValue Temp1 =
+ DAG.getNode(ISD::SUB, DL, XLenVT, EVL, DAG.getConstant(1, DL, XLenVT));
SDValue Temp2 = DAG.getNode(ISD::MUL, DL, XLenVT, Temp1,
DAG.getConstant(ElemWidthByte, DL, XLenVT));
- SDValue Base = DAG.getNode(ISD::ADD, DL, XLenVT, VPLoad->getBasePtr(), Temp2);
SDValue Stride = DAG.getSignedConstant(-ElemWidthByte, DL, XLenVT);
+ for (auto [VPLoad, LoadMask] : zip_equal(VPLoads, LoadMasks)) {
+ // Base = LoadAddr + (NumElem - 1) * ElemWidthByte
+ SDValue Base =
+ DAG.getNode(ISD::ADD, DL, XLenVT, VPLoad->getBasePtr(), Temp2);
- MachineFunction &MF = DAG.getMachineFunction();
- MachinePointerInfo PtrInfo(VPLoad->getAddressSpace());
- MachineMemOperand *MMO = MF.getMachineMemOperand(
- PtrInfo, VPLoad->getMemOperand()->getFlags(),
- LocationSize::beforeOrAfterPointer(), VPLoad->getAlign());
-
- SDValue Ret = DAG.getStridedLoadVP(
- LoadVT, DL, VPLoad->getChain(), Base, Stride, LoadMask,
- VPLoad->getVectorLength(), MMO, VPLoad->isExpandingLoad());
+ MachineFunction &MF = DAG.getMachineFunction();
+ MachinePointerInfo PtrInfo(VPLoad->getAddressSpace());
+ MachineMemOperand *MMO = MF.getMachineMemOperand(
+ PtrInfo, VPLoad->getMemOperand()->getFlags(),
+ LocationSize::beforeOrAfterPointer(), VPLoad->getAlign());
- DAG.ReplaceAllUsesOfValueWith(SDValue(VPLoad, 1), Ret.getValue(1));
+ SDValue Ret = DAG.getStridedLoadVP(
+ VPLoad->getValueType(0), DL, VPLoad->getChain(), Base, Stride, LoadMask,
+ VPLoad->getVectorLength(), MMO, VPLoad->isExpandingLoad());
+ DAG.ReplaceAllUsesWith(VPLoad, Ret.getNode());
+ }
- return Ret;
+ // Remove the top level reverse.
+ (void)sd_match(N, m_ReverseEVL(m_Value(Op), m_Specific(EVL)));
----------------
mshockwave wrote:
I don't think we need to use `m_Specific` here, as you only try to capture `Op` here
https://github.com/llvm/llvm-project/pull/205529
More information about the llvm-commits
mailing list