[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