[llvm] [X86] Fold OR of constant splats into GF2P8AFFINEQB (PR #194330)
via llvm-commits
llvm-commits at lists.llvm.org
Wed May 13 05:08:26 PDT 2026
================
@@ -53347,28 +53285,49 @@ static SDValue combineOrWithGF2P8AFFINEQB(SDNode *N, const SDLoc &DL,
using namespace SDPatternMatch;
assert(N->getOpcode() == ISD::OR && "Expected OR node");
- if (!N->getFlags().hasDisjoint() &&
- !DAG.haveNoCommonBitsSet(N->getOperand(0), N->getOperand(1)))
- return SDValue();
+ SDValue LHS = N->getOperand(0), RHS = N->getOperand(1);
- SDValue X, Y, SplatOp;
- APInt Imm, SplatVal;
+ SDValue X, Matrix, SplatOp;
+ APInt Imm;
+ if (!sd_match(N, m_Or(m_OneUse(m_TernaryOp(X86ISD::GF2P8AFFINEQB, m_Value(X),
+ m_Value(Matrix), m_ConstInt(Imm))),
+ m_Value(SplatOp))))
+ return SDValue();
- // Fold: (GF2P8AFFINEQB(X, Y, Imm) or_disjoint SplatVal)
- // -> GF2P8AFFINEQB(X, Y, Imm ^ SplatVal)
- // When OR is disjoint (no common bits), the splat constant can be folded
- // directly into the GF2P8AFFINEQB immediate via XOR.
+ APInt SplatVal;
+ if (!X86::isConstantSplat(SplatOp, SplatVal, /*AllowPartialUndefs=*/false))
+ return SDValue();
- if (sd_match(N, m_Or(m_OneUse(m_TernaryOp(X86ISD::GF2P8AFFINEQB, m_Value(X),
- m_Value(Y), m_ConstInt(Imm))),
- m_Value(SplatOp))) &&
- X86::isConstantSplat(SplatOp, SplatVal, /*AllowPartialUndefs=*/false)) {
+ if (N->getFlags().hasDisjoint() || DAG.haveNoCommonBitsSet(LHS, RHS)) {
+ // Fold: (GF2P8AFFINEQB(X, Matrix, Imm) or_disjoint SplatVal)
+ // -> GF2P8AFFINEQB(X, Matrix, Imm ^ SplatVal)
+ // When OR is disjoint (no common bits), the splat constant can be folded
+ // directly into the GF2P8AFFINEQB immediate via XOR.
uint64_t NewImm = (Imm.getZExtValue() ^ SplatVal.getZExtValue()) & 0xFF;
- return DAG.getNode(X86ISD::GF2P8AFFINEQB, DL, VT, X, Y,
+ return DAG.getNode(X86ISD::GF2P8AFFINEQB, DL, VT, X, Matrix,
DAG.getTargetConstant(NewImm, DL, MVT::i8));
}
- return SDValue();
+ APInt UndefElts;
+ SmallVector<APInt, 16> OldMatrix;
+ if (!getTargetConstantBitsFromNode(Matrix, 8, UndefElts, OldMatrix,
+ /*AllowWholeUndefs=*/false,
+ /*AllowPartialUndefs=*/false))
+ return SDValue();
+
+ uint8_t Mask8 = SplatVal.getZExtValue() & 0xFF;
+ uint8_t NewImm = Imm.getZExtValue() | Mask8;
+ SmallVector<SDValue, 64> MaskOps;
+ for (unsigned I = 0, E = VT.getVectorNumElements(); I != E; ++I) {
+ unsigned OutBit = 7 - (I & 7);
+ uint8_t Keep = ((Mask8 >> OutBit) & 1) ? 0x00 : 0xFF;
----------------
WalterKruger wrote:
```suggestion
uint8_t NewRow = ((Mask8 >> OutBit) & 1) ? 0x00 : OldMatrix[I];
```
It's cleaner to generate the new matrix directly as the build vector.
https://github.com/llvm/llvm-project/pull/194330
More information about the llvm-commits
mailing list