[SLP]Fix legality checks for bswap-based transformations Fix the checks for the non-power-of-2 base bswaps by checking the power-of-2 of the source type, not the target scalar type. Plus, add cost estimation for zext, if the source type does not match the scalar type. Fixes https://github.com/llvm/llvm-project/pull/184018#issuecomment-4053477562
diff --git a/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp b/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp index 6b459e0..8d09a1b 100644 --- a/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp +++ b/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
@@ -13267,8 +13267,6 @@ if (ScalarTy->isVectorTy()) return false; const unsigned Sz = DL->getTypeSizeInBits(ScalarTy); - if (!isPowerOf2_64(Sz)) - return false; const TreeEntry *LhsTE = getOperandEntry(&TE, /*Idx=*/0); const TreeEntry *RhsTE = getOperandEntry(&TE, /*Idx=*/1); // Lhs should be zext i<stride> to I<sz>. @@ -13280,7 +13278,8 @@ return false; Type *SrcScalarTy = cast<ZExtInst>(LhsTE->getMainOp())->getSrcTy(); unsigned Stride = DL->getTypeSizeInBits(SrcScalarTy); - if (!isPowerOf2_64(Stride) || Stride >= Sz) + if (!isPowerOf2_64(Stride) || Stride >= Sz || Sz % Stride != 0 || + !isPowerOf2_64(LhsTE->getVectorFactor())) return false; if (!(RhsTE->isGather() && RhsTE->ReorderIndices.empty() && RhsTE->ReuseShuffleIndices.empty() && !MinBWs.contains(RhsTE))) @@ -13332,6 +13331,8 @@ return false; } TTI::TargetCostKind CostKind = TTI::TCK_RecipThroughput; + auto *SrcType = IntegerType::getIntNTy(ScalarTy->getContext(), + Stride * LhsTE->getVectorFactor()); FastMathFlags FMF; SmallPtrSet<Value *, 4> CheckedExtracts; auto *VecTy = getWidenedType(ScalarTy, TE.getVectorFactor()); @@ -13347,7 +13348,7 @@ getWidenedType(SrcScalarTy, LhsTE->getVectorFactor()), CastCtx, CostKind); InstructionCost BitcastCost = TTI->getCastInstrCost( - Instruction::BitCast, ScalarTy, SrcVecTy, CastCtx, CostKind); + Instruction::BitCast, SrcType, SrcVecTy, CastCtx, CostKind); if (!Order.empty()) { fixupOrderingIndices(Order); SmallVector<int> Mask; @@ -13359,9 +13360,9 @@ constexpr unsigned ByteSize = 8; if (!Order.empty() && isReverseOrder(Order) && DL->getTypeSizeInBits(SrcScalarTy) == ByteSize) { - IntrinsicCostAttributes CostAttrs(Intrinsic::bswap, ScalarTy, {ScalarTy}); + IntrinsicCostAttributes CostAttrs(Intrinsic::bswap, SrcType, {SrcType}); InstructionCost BSwapCost = - TTI->getCastInstrCost(Instruction::BitCast, ScalarTy, SrcVecTy, CastCtx, + TTI->getCastInstrCost(Instruction::BitCast, SrcType, SrcVecTy, CastCtx, CostKind) + TTI->getIntrinsicInstrCost(CostAttrs, CostKind); if (BSwapCost <= BitcastCost) { @@ -13375,10 +13376,9 @@ SrcTE->getOpcode() == Instruction::Load && !SrcTE->isAltShuffle() && all_of(SrcTE->Scalars, [](Value *V) { return V->hasOneUse(); })) { auto *LI = cast<LoadInst>(SrcTE->getMainOp()); - IntrinsicCostAttributes CostAttrs(Intrinsic::bswap, ScalarTy, - {ScalarTy}); + IntrinsicCostAttributes CostAttrs(Intrinsic::bswap, SrcType, {SrcType}); InstructionCost BSwapCost = - TTI->getMemoryOpCost(Instruction::Load, ScalarTy, LI->getAlign(), + TTI->getMemoryOpCost(Instruction::Load, SrcType, LI->getAlign(), LI->getPointerAddressSpace(), CostKind) + TTI->getIntrinsicInstrCost(CostAttrs, CostKind); if (BSwapCost <= BitcastCost) { @@ -13399,7 +13399,7 @@ all_of(SrcTE->Scalars, [](Value *V) { return V->hasOneUse(); })) { auto *LI = cast<LoadInst>(SrcTE->getMainOp()); BitcastCost = - TTI->getMemoryOpCost(Instruction::Load, ScalarTy, LI->getAlign(), + TTI->getMemoryOpCost(Instruction::Load, SrcType, LI->getAlign(), LI->getPointerAddressSpace(), CostKind); VecCost += TTI->getMemoryOpCost(Instruction::Load, SrcVecTy, LI->getAlign(), @@ -13407,6 +13407,10 @@ ForLoads = true; } } + if (SrcType != ScalarTy) { + BitcastCost += TTI->getCastInstrCost(Instruction::ZExt, ScalarTy, SrcType, + TTI::CastContextHint::None, CostKind); + } return BitcastCost < VecCost; }
diff --git a/llvm/test/Transforms/SLPVectorizer/X86/non-power-of-2-bswap.ll b/llvm/test/Transforms/SLPVectorizer/X86/non-power-of-2-bswap.ll new file mode 100644 index 0000000..2f2b31b --- /dev/null +++ b/llvm/test/Transforms/SLPVectorizer/X86/non-power-of-2-bswap.ll
@@ -0,0 +1,26 @@ +; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --version 6 +; RUN: opt -passes=slp-vectorizer -S -slp-vectorize-non-power-of-2 -mtriple=x86_64-unknown-linux-gnu -mcpu=tigerlake < %s | FileCheck %s + +define i32 @test(i8 %0) { +; CHECK-LABEL: define i32 @test( +; CHECK-SAME: i8 [[TMP0:%.*]]) #[[ATTR0:[0-9]+]] { +; CHECK-NEXT: [[TMP2:%.*]] = insertelement <3 x i8> poison, i8 [[TMP0]], i32 0 +; CHECK-NEXT: [[TMP3:%.*]] = shufflevector <3 x i8> [[TMP2]], <3 x i8> poison, <3 x i32> zeroinitializer +; CHECK-NEXT: [[TMP4:%.*]] = and <3 x i8> [[TMP3]], splat (i8 1) +; CHECK-NEXT: [[TMP5:%.*]] = zext <3 x i8> [[TMP4]] to <3 x i32> +; CHECK-NEXT: [[TMP6:%.*]] = shl nuw nsw <3 x i32> [[TMP5]], <i32 16, i32 8, i32 0> +; CHECK-NEXT: [[TMP7:%.*]] = call i32 @llvm.vector.reduce.or.v3i32(<3 x i32> [[TMP6]]) +; CHECK-NEXT: ret i32 [[TMP7]] +; + %2 = and i8 %0, 1 + %3 = and i8 %0, 1 + %.sroa.5.0.insert.ext = zext nneg i8 %3 to i32 + %.sroa.5.0.insert.shift = shl nuw nsw i32 %.sroa.5.0.insert.ext, 16 + %.sroa.3.0.insert.ext = zext nneg i8 %2 to i32 + %.sroa.3.0.insert.shift = shl nuw nsw i32 %.sroa.3.0.insert.ext, 8 + %.sroa.3.0.insert.insert = or disjoint i32 %.sroa.5.0.insert.shift, %.sroa.3.0.insert.shift + %4 = and i8 %0, 1 + %.sroa.0.0.insert.ext = zext nneg i8 %4 to i32 + %.sroa.0.0.insert.insert = or disjoint i32 %.sroa.3.0.insert.insert, %.sroa.0.0.insert.ext + ret i32 %.sroa.0.0.insert.insert +}
diff --git a/llvm/test/Transforms/SLPVectorizer/non-power-of-2-bswap.ll b/llvm/test/Transforms/SLPVectorizer/non-power-of-2-bswap.ll index f0369e0..923663c 100644 --- a/llvm/test/Transforms/SLPVectorizer/non-power-of-2-bswap.ll +++ b/llvm/test/Transforms/SLPVectorizer/non-power-of-2-bswap.ll
@@ -4,24 +4,14 @@ define i64 @bswap_i24(ptr noalias %p, ptr noalias %p1) { ; CHECK-LABEL: define i64 @bswap_i24( ; CHECK-SAME: ptr noalias [[P:%.*]], ptr noalias [[P1:%.*]]) { -; CHECK-NEXT: [[G2:%.*]] = getelementptr i8, ptr [[P]], i32 2 -; CHECK-NEXT: [[T2:%.*]] = load i8, ptr [[G2]], align 1 -; CHECK-NEXT: [[G12:%.*]] = getelementptr i8, ptr [[P1]], i32 2 -; CHECK-NEXT: [[T12:%.*]] = load i8, ptr [[G12]], align 1 -; CHECK-NEXT: [[A2:%.*]] = add i8 [[T2]], [[T12]] -; CHECK-NEXT: [[Z2:%.*]] = zext i8 [[A2]] to i64 -; CHECK-NEXT: [[TMP1:%.*]] = load <2 x i8>, ptr [[P]], align 1 -; CHECK-NEXT: [[TMP2:%.*]] = load <2 x i8>, ptr [[P1]], align 1 -; CHECK-NEXT: [[TMP3:%.*]] = add <2 x i8> [[TMP1]], [[TMP2]] -; CHECK-NEXT: [[TMP4:%.*]] = zext <2 x i8> [[TMP3]] to <2 x i32> -; CHECK-NEXT: [[TMP5:%.*]] = shl <2 x i32> [[TMP4]], <i32 16, i32 8> -; CHECK-NEXT: [[TMP6:%.*]] = extractelement <2 x i32> [[TMP5]], i32 0 -; CHECK-NEXT: [[TMP7:%.*]] = zext i32 [[TMP6]] to i64 -; CHECK-NEXT: [[TMP8:%.*]] = extractelement <2 x i32> [[TMP5]], i32 1 +; CHECK-NEXT: [[TMP1:%.*]] = load <3 x i8>, ptr [[P]], align 1 +; CHECK-NEXT: [[TMP2:%.*]] = load <3 x i8>, ptr [[P1]], align 1 +; CHECK-NEXT: [[TMP3:%.*]] = add <3 x i8> [[TMP1]], [[TMP2]] +; CHECK-NEXT: [[TMP4:%.*]] = zext <3 x i8> [[TMP3]] to <3 x i32> +; CHECK-NEXT: [[TMP5:%.*]] = shl <3 x i32> [[TMP4]], <i32 16, i32 8, i32 0> +; CHECK-NEXT: [[TMP8:%.*]] = call i32 @llvm.vector.reduce.or.v3i32(<3 x i32> [[TMP5]]) ; CHECK-NEXT: [[TMP9:%.*]] = zext i32 [[TMP8]] to i64 -; CHECK-NEXT: [[OR01:%.*]] = or disjoint i64 [[TMP7]], [[TMP9]] -; CHECK-NEXT: [[OR012:%.*]] = or disjoint i64 [[OR01]], [[Z2]] -; CHECK-NEXT: ret i64 [[OR012]] +; CHECK-NEXT: ret i64 [[TMP9]] ; %g1 = getelementptr i8, ptr %p, i32 1 %g2 = getelementptr i8, ptr %p, i32 2