[llvm] [NVPTX][TTI] Update NVPTX TTI to enable f32x2 vectorization (PR #222994)

Daniel Donenfeld via llvm-commits llvm-commits at lists.llvm.org
Tue Sep 15 08:46:55 PDT 2026


https://github.com/daniel-donenfeld updated https://github.com/llvm/llvm-project/pull/222994

>From 8c6861fcf6adf920d95b2e15a52ce349fb2a37b8 Mon Sep 17 00:00:00 2001
From: Daniel Donenfeld <ddonenfeld at nvidia.com>
Date: Fri, 10 Jul 2026 15:48:57 +0000
Subject: [PATCH 1/2] [NVPTX][TTI] Update NVPTX TTI to enable f32x2
 vectorization

Enable f32x2 vectorization, and update the TTI cost model to better
model the newly exposed opportunities. This includes hooks to check for
native vector types, costing for non-native vectors, and costs for
inserting and extracting values from vectors. Add support for costing
splatted operands to f32x2 operations.

Also includes changes to enable folding vector add/mul into vector FMA
to better support creating f32x2 FMA operations.
---
 llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp     |   7 +-
 .../Target/NVPTX/NVPTXTargetTransformInfo.cpp | 249 +++++++++++++++++-
 .../Target/NVPTX/NVPTXTargetTransformInfo.h   |  62 ++---
 llvm/test/CodeGen/NVPTX/nvptx-fold-fma.ll     |  81 ++++++
 .../SLPVectorizer/NVPTX/f32-scalar-calls.ll   |  73 +++++
 .../Transforms/SLPVectorizer/NVPTX/f32x2.ll   | 201 ++++++++++++++
 .../SLPVectorizer/NVPTX/i8-dot-product.ll     | 150 +++++++++++
 .../SLPVectorizer/NVPTX/pair-copy-shapes.ll   |  39 +++
 .../SLPVectorizer/NVPTX/row-overhead.ll       | 225 ++++++++++++++++
 .../SLPVectorizer/NVPTX/v2i16-scalar-uses.ll  |  79 ++++++
 10 files changed, 1123 insertions(+), 43 deletions(-)
 create mode 100644 llvm/test/Transforms/SLPVectorizer/NVPTX/f32-scalar-calls.ll
 create mode 100644 llvm/test/Transforms/SLPVectorizer/NVPTX/f32x2.ll
 create mode 100644 llvm/test/Transforms/SLPVectorizer/NVPTX/i8-dot-product.ll
 create mode 100644 llvm/test/Transforms/SLPVectorizer/NVPTX/pair-copy-shapes.ll
 create mode 100644 llvm/test/Transforms/SLPVectorizer/NVPTX/row-overhead.ll
 create mode 100644 llvm/test/Transforms/SLPVectorizer/NVPTX/v2i16-scalar-uses.ll

diff --git a/llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp b/llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp
index bd16c7213b1e7..9db516928ff2c 100644
--- a/llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp
@@ -10,7 +10,7 @@
 // run late in the NVPTX IR pass pipeline just before the instruction selection.
 //
 // Currently, it implements the following transformation(s):
-// 1. FMA folding (float/double types):
+// 1. FMA folding (float/double types and vectors thereof):
 //    Transforms FMUL+FADD/FSUB sequences into FMA intrinsics when the
 //    'contract' fast-math flag is present. Supported patterns:
 //    - fadd(fmul(a, b), c) => fma(a, b, c)
@@ -125,8 +125,9 @@ static bool foldFMA(Function &F) {
       if (!BI->hasAllowContract())
         continue;
 
-      // Only float and double are supported.
-      if (!BI->getType()->isFloatTy() && !BI->getType()->isDoubleTy())
+      // Float, double, and vectors thereof are supported.
+      Type *ScalarTy = BI->getType()->getScalarType();
+      if (!ScalarTy->isFloatTy() && !ScalarTy->isDoubleTy())
         continue;
 
       if (tryFoldBinaryFMul(BI))
diff --git a/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.cpp b/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.cpp
index f52af84553648..bf0c43f32cea7 100644
--- a/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.cpp
@@ -462,10 +462,165 @@ NVPTXTTIImpl::getInstructionCost(const User *U,
   return BaseT::getInstructionCost(U, Operands, CostKind);
 }
 
+static bool isV2F32Ty(Type *Ty) {
+  // PTX has native packed arithmetic for pairs of f32 values.
+  auto *VTy = dyn_cast<FixedVectorType>(Ty);
+  return VTy && VTy->getNumElements() == 2 &&
+         VTy->getElementType()->isFloatTy();
+}
+
+static bool isF32VectorArithmeticOpcode(unsigned Opcode) {
+  // These opcodes map to packed f32x2 PTX arithmetic when the vector is v2f32.
+  switch (Opcode) {
+  case Instruction::FAdd:
+  case Instruction::FSub:
+  case Instruction::FMul:
+    return true;
+  default:
+    return false;
+  }
+}
+
+static bool isDemandedSplat(ArrayRef<Value *> VL, const APInt &DemandedElts) {
+  // Packed f32x2 arithmetic can consume scalar f32 operands as broadcasts, so
+  // splat buildvectors are free when the demanded lanes are identical.
+  if (VL.empty() || DemandedElts.getBitWidth() != VL.size())
+    return false;
+
+  Value *Splat = nullptr;
+  for (unsigned Idx = 0, E = VL.size(); Idx != E; ++Idx) {
+    if (!DemandedElts[Idx] || isa<UndefValue>(VL[Idx]))
+      continue;
+    if (!Splat) {
+      Splat = VL[Idx];
+      continue;
+    }
+    if (VL[Idx] != Splat)
+      return false;
+  }
+  return Splat != nullptr;
+}
+
+static bool isCheapPTXVectorInsertExtract(Type *Ty, const DataLayout &DL,
+                                          const TargetLoweringBase *TLI) {
+  assert(TLI && "Expected NVPTX TargetLowering");
+  auto *VTy = dyn_cast<FixedVectorType>(Ty);
+  return !VTy || NVPTX::isPackedVectorTy(TLI->getValueType(DL, Ty));
+}
+
+static bool hasNativeNVPTXVectorArithmetic(unsigned Opcode, Type *Ty,
+                                           const DataLayout &DL,
+                                           const TargetLoweringBase *TLI) {
+  assert(TLI && "Expected NVPTX TargetLowering");
+  if (!isa<FixedVectorType>(Ty))
+    return false;
+
+  if (isF32VectorArithmeticOpcode(Opcode) && isV2F32Ty(Ty))
+    return true;
+
+  int ISD = TLI->InstructionOpcodeToISD(Opcode);
+  if (ISD == 0)
+    return false;
+
+  EVT VT = TLI->getValueType(DL, Ty);
+  return TLI->isOperationLegal(ISD, VT);
+}
+
+static InstructionCost getVectorRegisterPieceCost(Type *Ty,
+                                                  const DataLayout &DL) {
+  auto *VTy = dyn_cast<FixedVectorType>(Ty);
+  if (!VTy)
+    return 0;
+
+  constexpr unsigned NVPTXRegBits = 32;
+  unsigned TyBits = DL.getTypeSizeInBits(Ty).getFixedValue();
+  unsigned NumPieces = divideCeil(TyBits, NVPTXRegBits);
+  return NumPieces;
+}
+
+static InstructionCost getNonNativeVectorOpPenalty(
+    Type *Ty, const DataLayout &DL, const TargetLoweringBase *TLI) {
+  auto *VTy = dyn_cast<FixedVectorType>(Ty);
+  if (!VTy || isCheapPTXVectorInsertExtract(Ty, DL, TLI))
+    return 0;
+
+  // Model a non-native vector operation by the scalar lane work plus the
+  // register-sized pieces that have to be assembled/disassembled around it.
+  return VTy->getNumElements() + getVectorRegisterPieceCost(Ty, DL);
+}
+
+static InstructionCost getNonNativeVectorShufflePenalty(
+    VectorType *DstTy, VectorType *SrcTy, VectorType *SubTp,
+    const DataLayout &DL, const TargetLoweringBase *TLI) {
+  InstructionCost Cost = 0;
+  for (VectorType *Ty : {DstTy, SrcTy, SubTp}) {
+    if (!Ty)
+      continue;
+    Cost += getNonNativeVectorOpPenalty(Ty, DL, TLI);
+  }
+  return Cost;
+}
+
+InstructionCost NVPTXTTIImpl::getShuffleCost(
+    TTI::ShuffleKind Kind, VectorType *DstTy, VectorType *SrcTy,
+    ArrayRef<int> Mask, TTI::TargetCostKind CostKind, int Index,
+    VectorType *SubTp, ArrayRef<const Value *> Args,
+    const Instruction *CxtI) const {
+  InstructionCost Cost =
+      BaseT::getShuffleCost(Kind, DstTy, SrcTy, Mask, CostKind, Index, SubTp,
+                            Args, CxtI);
+
+  if (CostKind != TTI::TCK_RecipThroughput)
+    return Cost;
+
+  // A scalar f32 broadcast feeding f32x2 arithmetic can use the scalar operand
+  // form of the packed PTX instruction, so do not charge a shuffle for it.
+  if (isV2F32Ty(DstTy) &&
+      (Kind == TTI::SK_Broadcast ||
+       (Kind == TTI::SK_PermuteSingleSrc &&
+        ShuffleVectorInst::isZeroEltSplatMask(Mask, Mask.size())))) {
+    return TTI::TCC_Free;
+  }
+
+  return Cost + getNonNativeVectorShufflePenalty(DstTy, SrcTy, SubTp, DL, TLI);
+}
+
+InstructionCost NVPTXTTIImpl::getVectorInstrCost(
+    unsigned Opcode, Type *Val, TTI::TargetCostKind CostKind, unsigned Index,
+    const Value *Op0, const Value *Op1, TTI::VectorInstrContext VIC) const {
+  InstructionCost Cost =
+      BaseT::getVectorInstrCost(Opcode, Val, CostKind, Index, Op0, Op1, VIC);
+
+  if (CostKind != TTI::TCK_RecipThroughput ||
+      (Opcode != Instruction::InsertElement &&
+       Opcode != Instruction::ExtractElement) ||
+      VIC != TTI::VectorInstrContext::None ||
+      isCheapPTXVectorInsertExtract(Val, DL, TLI))
+    return Cost;
+
+  return Cost + getNonNativeVectorOpPenalty(Val, DL, TLI);
+}
+
 InstructionCost NVPTXTTIImpl::getArithmeticInstrCost(
     unsigned Opcode, Type *Ty, TTI::TargetCostKind CostKind,
     TTI::OperandValueInfo Op1Info, TTI::OperandValueInfo Op2Info,
     ArrayRef<const Value *> Args, const Instruction *CxtI) const {
+  InstructionCost Cost =
+      BaseT::getArithmeticInstrCost(Opcode, Ty, CostKind, Op1Info, Op2Info);
+  if (CostKind == TTI::TCK_RecipThroughput) {
+    if (auto *VTy = dyn_cast<FixedVectorType>(Ty);
+        VTy && !hasNativeNVPTXVectorArithmetic(Opcode, Ty, DL, TLI)) {
+      InstructionCost ScalarCost =
+          VTy->getNumElements() *
+          BaseT::getArithmeticInstrCost(Opcode, VTy->getElementType(), CostKind,
+                                        Op1Info, Op2Info);
+      InstructionCost LegalizedCost =
+          ScalarCost + getNonNativeVectorOpPenalty(Ty, DL, TLI);
+      if (LegalizedCost > Cost)
+        Cost = LegalizedCost;
+    }
+  }
+
   // Legalize the type.
   std::pair<InstructionCost, MVT> LT = getTypeLegalizationCost(Ty);
 
@@ -473,8 +628,7 @@ InstructionCost NVPTXTTIImpl::getArithmeticInstrCost(
 
   switch (ISD) {
   default:
-    return BaseT::getArithmeticInstrCost(Opcode, Ty, CostKind, Op1Info,
-                                         Op2Info);
+    return Cost;
   case ISD::ADD:
   case ISD::MUL:
   case ISD::XOR:
@@ -486,9 +640,96 @@ InstructionCost NVPTXTTIImpl::getArithmeticInstrCost(
     if (LT.second.SimpleTy == MVT::i64)
       return 2 * LT.first;
     // Delegate other cases to the basic TTI.
-    return BaseT::getArithmeticInstrCost(Opcode, Ty, CostKind, Op1Info,
-                                         Op2Info);
+    return Cost;
+  }
+}
+
+InstructionCost NVPTXTTIImpl::getScalarizationOverhead(
+    VectorType *InTy, const APInt &DemandedElts, bool Insert, bool Extract,
+    TTI::TargetCostKind CostKind, bool ForPoisonSrc, ArrayRef<Value *> VL,
+    TTI::VectorInstrContext VIC) const {
+  if (!InTy->getElementCount().isFixed())
+    return InstructionCost::getInvalid();
+
+  auto VT = getTLI()->getValueType(DL, InTy);
+  auto NumElements = InTy->getElementCount().getFixedValue();
+  InstructionCost Cost = 0;
+  if (Insert && !VL.empty()) {
+    bool AllConstant = all_of(seq(NumElements), [&](int Idx) {
+      return !DemandedElts[Idx] || isa<Constant>(VL[Idx]);
+    });
+    if (AllConstant) {
+      Cost += TTI::TCC_Free;
+      Insert = false;
+    } else if (isV2F32Ty(InTy) && isDemandedSplat(VL, DemandedElts)) {
+      // A splat buildvector can be represented by a scalar broadcast operand of
+      // the packed f32 instruction.
+      Cost += TTI::TCC_Free;
+      Insert = false;
+    }
+  }
+  if (Insert && NVPTX::isPackedVectorTy(VT) && VT.is32BitVector()) {
+    // Can be built in a single 32-bit mov (64-bit regs are emulated in SASS
+    // with 2x 32-bit regs)
+    Cost += 1;
+    Insert = false;
+  }
+  if (Insert && VT == MVT::v4i8) {
+    Cost += 3; // 3 x PRMT
+    for (auto Idx : seq(NumElements))
+      if (DemandedElts[Idx])
+        Cost += 1; // zext operand to i32
+    Insert = false;
+  }
+  InstructionCost BaseCost = BaseT::getScalarizationOverhead(
+      InTy, DemandedElts, Insert, Extract, CostKind, ForPoisonSrc, VL, VIC);
+
+  if (Extract && CostKind == TTI::TCK_RecipThroughput &&
+      !NVPTX::isPackedVectorTy(VT) && !VT.is32BitVector() &&
+      !VT.is16BitVector()) {
+    InstructionCost ExtractCost = DemandedElts.popcount();
+    if (!Insert)
+      return Cost + ExtractCost;
+  }
+
+  return Cost + BaseCost;
+}
+
+InstructionCost NVPTXTTIImpl::getCastInstrCost(
+    unsigned Opcode, Type *Dst, Type *Src, TTI::CastContextHint CCH,
+    TTI::TargetCostKind CostKind, const Instruction *I) const {
+  InstructionCost Cost =
+      BaseT::getCastInstrCost(Opcode, Dst, Src, CCH, CostKind, I);
+  if (CostKind != TTI::TCK_RecipThroughput)
+    return Cost;
+
+  auto *DstVTy = dyn_cast<FixedVectorType>(Dst);
+  auto *SrcVTy = dyn_cast<FixedVectorType>(Src);
+  if (!DstVTy || !SrcVTy)
+    return Cost;
+
+  if (DstVTy->getNumElements() == SrcVTy->getNumElements()) {
+    InstructionCost ScalarCost =
+        DstVTy->getNumElements() *
+        BaseT::getCastInstrCost(Opcode, DstVTy->getElementType(),
+                                SrcVTy->getElementType(), CCH, CostKind, I);
+    InstructionCost LegalizedCost =
+        ScalarCost + getNonNativeVectorOpPenalty(Dst, DL, TLI) +
+        getNonNativeVectorOpPenalty(Src, DL, TLI);
+    if (LegalizedCost > Cost)
+      Cost = LegalizedCost;
+  } else if (Opcode == Instruction::BitCast) {
+    int ISD = TLI->InstructionOpcodeToISD(Opcode);
+    if (ISD == 0)
+      return Cost + getNonNativeVectorOpPenalty(Dst, DL, TLI) +
+             getNonNativeVectorOpPenalty(Src, DL, TLI);
+
+    EVT DstVT = TLI->getValueType(DL, Dst);
+    if (!TLI->isOperationLegal(ISD, DstVT))
+      Cost += getNonNativeVectorOpPenalty(Dst, DL, TLI) +
+              getNonNativeVectorOpPenalty(Src, DL, TLI);
   }
+  return Cost;
 }
 
 void NVPTXTTIImpl::getUnrollingPreferences(
diff --git a/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h b/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h
index 8bdafd6b905f1..6e586581b1384 100644
--- a/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h
+++ b/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h
@@ -42,6 +42,8 @@ class NVPTXTTIImpl final : public BasicTTIImplBase<NVPTXTTIImpl> {
   bool isSourceOfDivergence(const Value *V) const;
 
 public:
+  using BaseT::getVectorInstrCost;
+
   explicit NVPTXTTIImpl(const NVPTXTargetMachine *TM, const Function &F)
       : BaseT(TM, F.getDataLayout()), ST(TM->getSubtargetImpl()),
         TLI(ST->getTargetLowering()) {}
@@ -82,11 +84,11 @@ class NVPTXTTIImpl final : public BasicTTIImplBase<NVPTXTTIImpl> {
   // LoopVectorizer's unrolling heuristics.
   unsigned getNumberOfRegisters(unsigned ClassID) const override { return 1; }
 
-  // Only <2 x half> should be vectorized, so always return 32 for the vector
-  // register size.
+  // The types  <2 x half> and  <2 x float> can be vectorized, so return 64 for
+  // the vector register size.
   TypeSize
   getRegisterBitWidth(TargetTransformInfo::RegisterKind K) const override {
-    return TypeSize::getFixed(32);
+    return TypeSize::getFixed(64);
   }
   unsigned getMinVectorRegisterBitWidth() const override { return 32; }
 
@@ -120,45 +122,33 @@ class NVPTXTTIImpl final : public BasicTTIImplBase<NVPTXTTIImpl> {
       ArrayRef<const Value *> Args = {},
       const Instruction *CxtI = nullptr) const override;
 
+  InstructionCost
+  getCastInstrCost(unsigned Opcode, Type *Dst, Type *Src,
+                   TTI::CastContextHint CCH, TTI::TargetCostKind CostKind,
+                   const Instruction *I = nullptr) const override;
+
+  InstructionCost
+  getShuffleCost(TTI::ShuffleKind Kind, VectorType *DstTy, VectorType *SrcTy,
+                 ArrayRef<int> Mask = {},
+                 TTI::TargetCostKind CostKind = TTI::TCK_RecipThroughput,
+                 int Index = 0, VectorType *SubTp = nullptr,
+                 ArrayRef<const Value *> Args = {},
+                 const Instruction *CxtI = nullptr) const override;
+
+  InstructionCost
+  getVectorInstrCost(unsigned Opcode, Type *Val, TTI::TargetCostKind CostKind,
+                     unsigned Index = -1, const Value *Op0 = nullptr,
+                     const Value *Op1 = nullptr,
+                     TTI::VectorInstrContext VIC =
+                         TTI::VectorInstrContext::None) const override;
+
   InstructionCost
   getScalarizationOverhead(VectorType *InTy, const APInt &DemandedElts,
                            bool Insert, bool Extract,
                            TTI::TargetCostKind CostKind,
                            bool ForPoisonSrc = true, ArrayRef<Value *> VL = {},
                            TTI::VectorInstrContext VIC =
-                               TTI::VectorInstrContext::None) const override {
-    if (!InTy->getElementCount().isFixed())
-      return InstructionCost::getInvalid();
-
-    auto VT = getTLI()->getValueType(DL, InTy);
-    auto NumElements = InTy->getElementCount().getFixedValue();
-    InstructionCost Cost = 0;
-    if (Insert && !VL.empty()) {
-      bool AllConstant = all_of(seq(NumElements), [&](int Idx) {
-        return !DemandedElts[Idx] || isa<Constant>(VL[Idx]);
-      });
-      if (AllConstant) {
-        Cost += TTI::TCC_Free;
-        Insert = false;
-      }
-    }
-    if (Insert && NVPTX::isPackedVectorTy(VT) && VT.is32BitVector()) {
-      // Can be built in a single 32-bit mov (64-bit regs are emulated in SASS
-      // with 2x 32-bit regs)
-      Cost += 1;
-      Insert = false;
-    }
-    if (Insert && VT == MVT::v4i8) {
-      InstructionCost Cost = 3; // 3 x PRMT
-      for (auto Idx : seq(NumElements))
-        if (DemandedElts[Idx])
-          Cost += 1; // zext operand to i32
-      Insert = false;
-    }
-    return Cost + BaseT::getScalarizationOverhead(InTy, DemandedElts, Insert,
-                                                  Extract, CostKind,
-                                                  ForPoisonSrc, VL);
-  }
+                               TTI::VectorInstrContext::None) const override;
 
   void getUnrollingPreferences(Loop *L, ScalarEvolution &SE,
                                TTI::UnrollingPreferences &UP,
diff --git a/llvm/test/CodeGen/NVPTX/nvptx-fold-fma.ll b/llvm/test/CodeGen/NVPTX/nvptx-fold-fma.ll
index 6d9ad8d3ad436..6c2cf8e5afd84 100644
--- a/llvm/test/CodeGen/NVPTX/nvptx-fold-fma.ll
+++ b/llvm/test/CodeGen/NVPTX/nvptx-fold-fma.ll
@@ -245,3 +245,84 @@ define double @test_fadd_fmul_c_double(double %a, double %b, double %c) {
   %add = fadd contract double %mul, %c
   ret double %add
 }
+
+
+; fadd(fmul(a, b), c) => fma(a, b, c)
+define <2 x float> @test_fadd_fmul_c_v2f32(<2 x float> %a, <2 x float> %b, <2 x float> %c) {
+; CHECK-LABEL: define <2 x float> @test_fadd_fmul_c_v2f32(
+; CHECK-SAME: <2 x float> [[A:%.*]], <2 x float> [[B:%.*]], <2 x float> [[C:%.*]]) {
+; CHECK-NEXT:    [[ADD:%.*]] = call contract <2 x float> @llvm.fma.v2f32(<2 x float> [[A]], <2 x float> [[B]], <2 x float> [[C]])
+; CHECK-NEXT:    ret <2 x float> [[ADD]]
+;
+  %mul = fmul contract <2 x float> %a, %b
+  %add = fadd contract <2 x float> %mul, %c
+  ret <2 x float> %add
+}
+
+
+; fadd(c, fmul(a, b)) => fma(a, b, c)
+define <2 x float> @test_fadd_c_fmul_v2f32(<2 x float> %a, <2 x float> %b, <2 x float> %c) {
+; CHECK-LABEL: define <2 x float> @test_fadd_c_fmul_v2f32(
+; CHECK-SAME: <2 x float> [[A:%.*]], <2 x float> [[B:%.*]], <2 x float> [[C:%.*]]) {
+; CHECK-NEXT:    [[ADD:%.*]] = call contract <2 x float> @llvm.fma.v2f32(<2 x float> [[A]], <2 x float> [[B]], <2 x float> [[C]])
+; CHECK-NEXT:    ret <2 x float> [[ADD]]
+;
+  %mul = fmul contract <2 x float> %a, %b
+  %add = fadd contract <2 x float> %c, %mul
+  ret <2 x float> %add
+}
+
+
+; fsub(fmul(a, b), c) => fma(a, b, fneg(c))
+define <2 x float> @test_fsub_fmul_c_v2f32(<2 x float> %a, <2 x float> %b, <2 x float> %c) {
+; CHECK-LABEL: define <2 x float> @test_fsub_fmul_c_v2f32(
+; CHECK-SAME: <2 x float> [[A:%.*]], <2 x float> [[B:%.*]], <2 x float> [[C:%.*]]) {
+; CHECK-NEXT:    [[TMP1:%.*]] = fneg contract <2 x float> [[C]]
+; CHECK-NEXT:    [[SUB:%.*]] = call contract <2 x float> @llvm.fma.v2f32(<2 x float> [[A]], <2 x float> [[B]], <2 x float> [[TMP1]])
+; CHECK-NEXT:    ret <2 x float> [[SUB]]
+;
+  %mul = fmul contract <2 x float> %a, %b
+  %sub = fsub contract <2 x float> %mul, %c
+  ret <2 x float> %sub
+}
+
+
+; fsub(c, fmul(a, b)) => fma(fneg(a), b, c)
+define <2 x float> @test_fsub_c_fmul_v2f32(<2 x float> %a, <2 x float> %b, <2 x float> %c) {
+; CHECK-LABEL: define <2 x float> @test_fsub_c_fmul_v2f32(
+; CHECK-SAME: <2 x float> [[A:%.*]], <2 x float> [[B:%.*]], <2 x float> [[C:%.*]]) {
+; CHECK-NEXT:    [[TMP1:%.*]] = fneg contract <2 x float> [[A]]
+; CHECK-NEXT:    [[SUB:%.*]] = call contract <2 x float> @llvm.fma.v2f32(<2 x float> [[TMP1]], <2 x float> [[B]], <2 x float> [[C]])
+; CHECK-NEXT:    ret <2 x float> [[SUB]]
+;
+  %mul = fmul contract <2 x float> %a, %b
+  %sub = fsub contract <2 x float> %c, %mul
+  ret <2 x float> %sub
+}
+
+
+; fadd(fmul(a, b), c) => fma(a, b, c)
+define <2 x double> @test_fadd_fmul_c_v2f64(<2 x double> %a, <2 x double> %b, <2 x double> %c) {
+; CHECK-LABEL: define <2 x double> @test_fadd_fmul_c_v2f64(
+; CHECK-SAME: <2 x double> [[A:%.*]], <2 x double> [[B:%.*]], <2 x double> [[C:%.*]]) {
+; CHECK-NEXT:    [[ADD:%.*]] = call contract <2 x double> @llvm.fma.v2f64(<2 x double> [[A]], <2 x double> [[B]], <2 x double> [[C]])
+; CHECK-NEXT:    ret <2 x double> [[ADD]]
+;
+  %mul = fmul contract <2 x double> %a, %b
+  %add = fadd contract <2 x double> %mul, %c
+  ret <2 x double> %add
+}
+
+
+; fsub(fmul(a, b), c) => fma(a, b, fneg(c))
+define <2 x double> @test_fsub_fmul_c_v2f64(<2 x double> %a, <2 x double> %b, <2 x double> %c) {
+; CHECK-LABEL: define <2 x double> @test_fsub_fmul_c_v2f64(
+; CHECK-SAME: <2 x double> [[A:%.*]], <2 x double> [[B:%.*]], <2 x double> [[C:%.*]]) {
+; CHECK-NEXT:    [[TMP1:%.*]] = fneg contract <2 x double> [[C]]
+; CHECK-NEXT:    [[SUB:%.*]] = call contract <2 x double> @llvm.fma.v2f64(<2 x double> [[A]], <2 x double> [[B]], <2 x double> [[TMP1]])
+; CHECK-NEXT:    ret <2 x double> [[SUB]]
+;
+  %mul = fmul contract <2 x double> %a, %b
+  %sub = fsub contract <2 x double> %mul, %c
+  ret <2 x double> %sub
+}
diff --git a/llvm/test/Transforms/SLPVectorizer/NVPTX/f32-scalar-calls.ll b/llvm/test/Transforms/SLPVectorizer/NVPTX/f32-scalar-calls.ll
new file mode 100644
index 0000000000000..a7d0ea63732ea
--- /dev/null
+++ b/llvm/test/Transforms/SLPVectorizer/NVPTX/f32-scalar-calls.ll
@@ -0,0 +1,73 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py
+; RUN: opt < %s -passes=slp-vectorizer -S -mtriple=nvptx64-nvidia-cuda -mcpu=sm_100 | FileCheck %s
+
+target triple = "nvptx64-nvidia-cuda"
+
+declare float @use_f32(float, float, i32)
+declare i32 @use_i32(float, i32)
+
+; Adjacent f32 values selected through a branch should not be packed when all
+; uses need scalar call arguments.
+define void @f32_pair_scalar_calls(ptr %out, ptr %outi, ptr %a, ptr %b, i32 %p) {
+; CHECK-LABEL: @f32_pair_scalar_calls(
+; CHECK-NEXT:  entry:
+; CHECK-NEXT:    [[COND:%.*]] = and i32 [[P:%.*]], 1
+; CHECK-NEXT:    [[TOBOOL:%.*]] = icmp eq i32 [[COND]], 0
+; CHECK-NEXT:    br i1 [[TOBOOL]], label [[ELSE:%.*]], label [[THEN:%.*]]
+; CHECK:       then:
+; CHECK-NEXT:    [[A0:%.*]] = load float, ptr [[A:%.*]], align 8
+; CHECK-NEXT:    [[A1P:%.*]] = getelementptr float, ptr [[A]], i64 1
+; CHECK-NEXT:    [[A1:%.*]] = load float, ptr [[A1P]], align 4
+; CHECK-NEXT:    br label [[JOIN:%.*]]
+; CHECK:       else:
+; CHECK-NEXT:    [[B0:%.*]] = load float, ptr [[B:%.*]], align 8
+; CHECK-NEXT:    [[B1P:%.*]] = getelementptr float, ptr [[B]], i64 1
+; CHECK-NEXT:    [[B1:%.*]] = load float, ptr [[B1P]], align 4
+; CHECK-NEXT:    br label [[JOIN]]
+; CHECK:       join:
+; CHECK-NEXT:    [[X0:%.*]] = phi float [ [[A0]], [[THEN]] ], [ [[B0]], [[ELSE]] ]
+; CHECK-NEXT:    [[X1:%.*]] = phi float [ [[A1]], [[THEN]] ], [ [[B1]], [[ELSE]] ]
+; CHECK-NEXT:    [[C0:%.*]] = call float @use_f32(float [[X0]], float [[X1]], i32 [[P]])
+; CHECK-NEXT:    [[C1:%.*]] = call float @use_f32(float [[X1]], float [[X0]], i32 [[P]])
+; CHECK-NEXT:    [[I0:%.*]] = call i32 @use_i32(float [[X0]], i32 [[P]])
+; CHECK-NEXT:    [[I1:%.*]] = call i32 @use_i32(float [[X1]], i32 [[P]])
+; CHECK-NEXT:    store float [[C0]], ptr [[OUT:%.*]], align 8
+; CHECK-NEXT:    [[OUT1:%.*]] = getelementptr float, ptr [[OUT]], i64 1
+; CHECK-NEXT:    store float [[C1]], ptr [[OUT1]], align 4
+; CHECK-NEXT:    store i32 [[I0]], ptr [[OUTI:%.*]], align 8
+; CHECK-NEXT:    [[OUTI1:%.*]] = getelementptr i32, ptr [[OUTI]], i64 1
+; CHECK-NEXT:    store i32 [[I1]], ptr [[OUTI1]], align 4
+; CHECK-NEXT:    ret void
+;
+entry:
+  %cond = and i32 %p, 1
+  %tobool = icmp eq i32 %cond, 0
+  br i1 %tobool, label %else, label %then
+
+then:
+  %a0 = load float, ptr %a, align 8
+  %a1p = getelementptr float, ptr %a, i64 1
+  %a1 = load float, ptr %a1p, align 4
+  br label %join
+
+else:
+  %b0 = load float, ptr %b, align 8
+  %b1p = getelementptr float, ptr %b, i64 1
+  %b1 = load float, ptr %b1p, align 4
+  br label %join
+
+join:
+  %x0 = phi float [ %a0, %then ], [ %b0, %else ]
+  %x1 = phi float [ %a1, %then ], [ %b1, %else ]
+  %c0 = call float @use_f32(float %x0, float %x1, i32 %p)
+  %c1 = call float @use_f32(float %x1, float %x0, i32 %p)
+  %i0 = call i32 @use_i32(float %x0, i32 %p)
+  %i1 = call i32 @use_i32(float %x1, i32 %p)
+  store float %c0, ptr %out, align 8
+  %out1 = getelementptr float, ptr %out, i64 1
+  store float %c1, ptr %out1, align 4
+  store i32 %i0, ptr %outi, align 8
+  %outi1 = getelementptr i32, ptr %outi, i64 1
+  store i32 %i1, ptr %outi1, align 4
+  ret void
+}
diff --git a/llvm/test/Transforms/SLPVectorizer/NVPTX/f32x2.ll b/llvm/test/Transforms/SLPVectorizer/NVPTX/f32x2.ll
new file mode 100644
index 0000000000000..1b852fc66d17b
--- /dev/null
+++ b/llvm/test/Transforms/SLPVectorizer/NVPTX/f32x2.ll
@@ -0,0 +1,201 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py
+; RUN: opt < %s -passes=slp-vectorizer -S -mtriple=nvptx64-nvidia-cuda -mcpu=sm_100 | FileCheck %s
+
+target triple = "nvptx64-nvidia-cuda"
+
+define void @pair_f32_add(ptr %out, ptr %a, ptr %b) {
+; CHECK-LABEL: @pair_f32_add(
+; CHECK-NEXT:  entry:
+; CHECK-NEXT:    [[TMP0:%.*]] = load <2 x float>, ptr [[A:%.*]], align 8
+; CHECK-NEXT:    [[TMP1:%.*]] = load <2 x float>, ptr [[B:%.*]], align 8
+; CHECK-NEXT:    [[TMP2:%.*]] = fadd <2 x float> [[TMP0]], [[TMP1]]
+; CHECK-NEXT:    store <2 x float> [[TMP2]], ptr [[OUT:%.*]], align 8
+; CHECK-NEXT:    ret void
+;
+entry:
+  %a0 = load float, ptr %a, align 8
+  %a1p = getelementptr float, ptr %a, i64 1
+  %a1 = load float, ptr %a1p, align 4
+  %b0 = load float, ptr %b, align 8
+  %b1p = getelementptr float, ptr %b, i64 1
+  %b1 = load float, ptr %b1p, align 4
+  %r0 = fadd float %a0, %b0
+  %r1 = fadd float %a1, %b1
+  store float %r0, ptr %out, align 8
+  %out1 = getelementptr float, ptr %out, i64 1
+  store float %r1, ptr %out1, align 4
+  ret void
+}
+
+define void @pair_f32_sub(ptr %out, ptr %a, ptr %b) {
+; CHECK-LABEL: @pair_f32_sub(
+; CHECK-NEXT:  entry:
+; CHECK-NEXT:    [[TMP0:%.*]] = load <2 x float>, ptr [[A:%.*]], align 8
+; CHECK-NEXT:    [[TMP1:%.*]] = load <2 x float>, ptr [[B:%.*]], align 8
+; CHECK-NEXT:    [[TMP2:%.*]] = fsub <2 x float> [[TMP0]], [[TMP1]]
+; CHECK-NEXT:    store <2 x float> [[TMP2]], ptr [[OUT:%.*]], align 8
+; CHECK-NEXT:    ret void
+;
+entry:
+  %a0 = load float, ptr %a, align 8
+  %a1p = getelementptr float, ptr %a, i64 1
+  %a1 = load float, ptr %a1p, align 4
+  %b0 = load float, ptr %b, align 8
+  %b1p = getelementptr float, ptr %b, i64 1
+  %b1 = load float, ptr %b1p, align 4
+  %r0 = fsub float %a0, %b0
+  %r1 = fsub float %a1, %b1
+  store float %r0, ptr %out, align 8
+  %out1 = getelementptr float, ptr %out, i64 1
+  store float %r1, ptr %out1, align 4
+  ret void
+}
+
+define void @pair_f32_mul(ptr %out, ptr %a, ptr %b) {
+; CHECK-LABEL: @pair_f32_mul(
+; CHECK-NEXT:  entry:
+; CHECK-NEXT:    [[TMP0:%.*]] = load <2 x float>, ptr [[A:%.*]], align 8
+; CHECK-NEXT:    [[TMP1:%.*]] = load <2 x float>, ptr [[B:%.*]], align 8
+; CHECK-NEXT:    [[TMP2:%.*]] = fmul <2 x float> [[TMP0]], [[TMP1]]
+; CHECK-NEXT:    store <2 x float> [[TMP2]], ptr [[OUT:%.*]], align 8
+; CHECK-NEXT:    ret void
+;
+entry:
+  %a0 = load float, ptr %a, align 8
+  %a1p = getelementptr float, ptr %a, i64 1
+  %a1 = load float, ptr %a1p, align 4
+  %b0 = load float, ptr %b, align 8
+  %b1p = getelementptr float, ptr %b, i64 1
+  %b1 = load float, ptr %b1p, align 4
+  %r0 = fmul float %a0, %b0
+  %r1 = fmul float %a1, %b1
+  store float %r0, ptr %out, align 8
+  %out1 = getelementptr float, ptr %out, i64 1
+  store float %r1, ptr %out1, align 4
+  ret void
+}
+
+define void @pair_f32_splat(ptr %out, ptr %a, float %x) {
+; CHECK-LABEL: @pair_f32_splat(
+; CHECK-NEXT:  entry:
+; CHECK-NEXT:    [[TMP0:%.*]] = load <2 x float>, ptr [[A:%.*]], align 8
+; CHECK-NEXT:    [[TMP1:%.*]] = insertelement <2 x float> poison, float [[X:%.*]], i64 0
+; CHECK-NEXT:    [[TMP2:%.*]] = shufflevector <2 x float> [[TMP1]], <2 x float> poison, <2 x i32> zeroinitializer
+; CHECK-NEXT:    [[TMP3:%.*]] = fmul <2 x float> [[TMP0]], [[TMP2]]
+; CHECK-NEXT:    store <2 x float> [[TMP3]], ptr [[OUT:%.*]], align 8
+; CHECK-NEXT:    ret void
+;
+entry:
+  %a0 = load float, ptr %a, align 8
+  %a1p = getelementptr float, ptr %a, i64 1
+  %a1 = load float, ptr %a1p, align 4
+  %r0 = fmul float %a0, %x
+  %r1 = fmul float %a1, %x
+  store float %r0, ptr %out, align 8
+  %out1 = getelementptr float, ptr %out, i64 1
+  store float %r1, ptr %out1, align 4
+  ret void
+}
+
+define void @pair_f32_constant_splat(ptr %out, ptr %a) {
+; CHECK-LABEL: @pair_f32_constant_splat(
+; CHECK-NEXT:  entry:
+; CHECK-NEXT:    [[TMP0:%.*]] = load <2 x float>, ptr [[A:%.*]], align 8
+; CHECK-NEXT:    [[TMP1:%.*]] = fadd <2 x float> [[TMP0]], splat (float 4.000000e+00)
+; CHECK-NEXT:    store <2 x float> [[TMP1]], ptr [[OUT:%.*]], align 8
+; CHECK-NEXT:    ret void
+;
+entry:
+  %a0 = load float, ptr %a, align 8
+  %a1p = getelementptr float, ptr %a, i64 1
+  %a1 = load float, ptr %a1p, align 4
+  %r0 = fadd float %a0, 4.000000e+00
+  %r1 = fadd float %a1, 4.000000e+00
+  store float %r0, ptr %out, align 8
+  %out1 = getelementptr float, ptr %out, i64 1
+  store float %r1, ptr %out1, align 4
+  ret void
+}
+
+define void @pair_f32_loaded_splat_used(ptr %out, ptr %a, ptr %xp) {
+; CHECK-LABEL: @pair_f32_loaded_splat_used(
+; CHECK-NEXT:  entry:
+; CHECK-NEXT:    [[X:%.*]] = load float, ptr [[XP:%.*]], align 4
+; CHECK-NEXT:    [[TMP0:%.*]] = load <2 x float>, ptr [[A:%.*]], align 8
+; CHECK-NEXT:    [[TMP1:%.*]] = insertelement <2 x float> poison, float [[X]], i64 0
+; CHECK-NEXT:    [[TMP2:%.*]] = shufflevector <2 x float> [[TMP1]], <2 x float> poison, <2 x i32> zeroinitializer
+; CHECK-NEXT:    [[TMP3:%.*]] = fmul <2 x float> [[TMP0]], [[TMP2]]
+; CHECK-NEXT:    store <2 x float> [[TMP3]], ptr [[OUT:%.*]], align 8
+; CHECK-NEXT:    [[OUT2:%.*]] = getelementptr float, ptr [[OUT]], i64 2
+; CHECK-NEXT:    store float [[X]], ptr [[OUT2]], align 4
+; CHECK-NEXT:    ret void
+;
+entry:
+  %x = load float, ptr %xp, align 4
+  %a0 = load float, ptr %a, align 8
+  %a1p = getelementptr float, ptr %a, i64 1
+  %a1 = load float, ptr %a1p, align 4
+  %r0 = fmul float %a0, %x
+  %r1 = fmul float %a1, %x
+  store float %r0, ptr %out, align 8
+  %out1 = getelementptr float, ptr %out, i64 1
+  store float %r1, ptr %out1, align 4
+  %out2 = getelementptr float, ptr %out, i64 2
+  store float %x, ptr %out2, align 4
+  ret void
+}
+
+
+; SLP should rebuild the scalar phi/insert chain as a packed f32x2 value.
+define void @f32x2_phi_store(ptr addrspace(1) %x, ptr addrspace(3) %scratch, i1 %cond) {
+; CHECK-LABEL: @f32x2_phi_store(
+; CHECK:       then:
+; CHECK-NEXT:    [[V:%.*]] = load <2 x float>, ptr addrspace(1) [[X:%.*]], align 8
+; CHECK-NEXT:    br label [[JOIN:%.*]]
+; CHECK:       join:
+; CHECK-NEXT:    [[PHI:%.*]] = phi <2 x float> [ [[V]], [[THEN:%.*]] ], [ zeroinitializer, [[ENTRY:%.*]] ]
+; CHECK-NEXT:    store <2 x float> [[PHI]], ptr addrspace(3) [[SCRATCH:%.*]], align 8
+; CHECK-NEXT:    ret void
+;
+entry:
+  br i1 %cond, label %then, label %join
+
+then:
+  %v = load <2 x float>, ptr addrspace(1) %x, align 8
+  %v0 = extractelement <2 x float> %v, i32 0
+  %v1 = extractelement <2 x float> %v, i32 1
+  br label %join
+
+join:
+  %p0 = phi float [ %v0, %then ], [ 0.0, %entry ]
+  %p1 = phi float [ %v1, %then ], [ 0.0, %entry ]
+  %i0 = insertelement <2 x float> poison, float %p0, i32 0
+  %i1 = insertelement <2 x float> %i0, float %p1, i32 1
+  store <2 x float> %i1, ptr addrspace(3) %scratch, align 8
+  ret void
+}
+
+define void @f32x2_add_store(ptr addrspace(3) %a, ptr addrspace(3) %b,
+                                   ptr addrspace(3) %out) {
+; CHECK-LABEL: @f32x2_add_store(
+; CHECK-NEXT:  entry:
+; CHECK-NEXT:    [[VA:%.*]] = load <2 x float>, ptr addrspace(3) [[A:%.*]], align 8
+; CHECK-NEXT:    [[VB:%.*]] = load <2 x float>, ptr addrspace(3) [[B:%.*]], align 8
+; CHECK-NEXT:    [[TMP1:%.*]] = fadd <2 x float> [[VA]], [[VB]]
+; CHECK-NEXT:    store <2 x float> [[TMP1]], ptr addrspace(3) [[OUT:%.*]], align 8
+; CHECK-NEXT:    ret void
+;
+entry:
+  %va = load <2 x float>, ptr addrspace(3) %a, align 8
+  %vb = load <2 x float>, ptr addrspace(3) %b, align 8
+  %a0 = extractelement <2 x float> %va, i32 0
+  %a1 = extractelement <2 x float> %va, i32 1
+  %b0 = extractelement <2 x float> %vb, i32 0
+  %b1 = extractelement <2 x float> %vb, i32 1
+  %s0 = fadd float %a0, %b0
+  %s1 = fadd float %a1, %b1
+  %i0 = insertelement <2 x float> poison, float %s0, i32 0
+  %i1 = insertelement <2 x float> %i0, float %s1, i32 1
+  store <2 x float> %i1, ptr addrspace(3) %out, align 8
+  ret void
+}
diff --git a/llvm/test/Transforms/SLPVectorizer/NVPTX/i8-dot-product.ll b/llvm/test/Transforms/SLPVectorizer/NVPTX/i8-dot-product.ll
new file mode 100644
index 0000000000000..cd718af0b0ef9
--- /dev/null
+++ b/llvm/test/Transforms/SLPVectorizer/NVPTX/i8-dot-product.ll
@@ -0,0 +1,150 @@
+; RUN: opt < %s -passes=slp-vectorizer -S -mtriple=nvptx64-nvidia-cuda -mcpu=sm_100 | FileCheck %s
+
+target triple = "nvptx64-nvidia-cuda"
+
+; Keep this scalar so later NVPTX lowering can recognize the i8 dot-product
+; idiom instead of unpacking v2i8/v2i16 operations.
+define i32 @i8-dot-product(ptr addrspace(3) %a, ptr addrspace(3) %b, i32 %acc) {
+; CHECK-LABEL: @i8-dot-product(
+; CHECK-NEXT:  entry:
+; CHECK-NEXT:    [[VA:%.*]] = load <8 x i8>, ptr addrspace(3) [[A:%.*]], align 8
+; CHECK-NEXT:    [[VB:%.*]] = load <8 x i8>, ptr addrspace(3) [[B:%.*]], align 8
+; CHECK-NEXT:    [[A0:%.*]] = extractelement <8 x i8> [[VA]], i32 0
+; CHECK-NEXT:    [[A1:%.*]] = extractelement <8 x i8> [[VA]], i32 1
+; CHECK-NEXT:    [[A2:%.*]] = extractelement <8 x i8> [[VA]], i32 2
+; CHECK-NEXT:    [[A3:%.*]] = extractelement <8 x i8> [[VA]], i32 3
+; CHECK-NEXT:    [[A4:%.*]] = extractelement <8 x i8> [[VA]], i32 4
+; CHECK-NEXT:    [[A5:%.*]] = extractelement <8 x i8> [[VA]], i32 5
+; CHECK-NEXT:    [[A6:%.*]] = extractelement <8 x i8> [[VA]], i32 6
+; CHECK-NEXT:    [[A7:%.*]] = extractelement <8 x i8> [[VA]], i32 7
+; CHECK-NEXT:    [[B0:%.*]] = extractelement <8 x i8> [[VB]], i32 0
+; CHECK-NEXT:    [[B1:%.*]] = extractelement <8 x i8> [[VB]], i32 1
+; CHECK-NEXT:    [[B2:%.*]] = extractelement <8 x i8> [[VB]], i32 2
+; CHECK-NEXT:    [[B3:%.*]] = extractelement <8 x i8> [[VB]], i32 3
+; CHECK-NEXT:    [[B4:%.*]] = extractelement <8 x i8> [[VB]], i32 4
+; CHECK-NEXT:    [[B5:%.*]] = extractelement <8 x i8> [[VB]], i32 5
+; CHECK-NEXT:    [[B6:%.*]] = extractelement <8 x i8> [[VB]], i32 6
+; CHECK-NEXT:    [[B7:%.*]] = extractelement <8 x i8> [[VB]], i32 7
+; CHECK-NEXT:    [[SA0:%.*]] = sext i8 [[A0]] to i32
+; CHECK-NEXT:    [[SA1:%.*]] = sext i8 [[A1]] to i32
+; CHECK-NEXT:    [[SA2:%.*]] = sext i8 [[A2]] to i32
+; CHECK-NEXT:    [[SA3:%.*]] = sext i8 [[A3]] to i32
+; CHECK-NEXT:    [[SA4:%.*]] = sext i8 [[A4]] to i32
+; CHECK-NEXT:    [[SA5:%.*]] = sext i8 [[A5]] to i32
+; CHECK-NEXT:    [[SA6:%.*]] = sext i8 [[A6]] to i32
+; CHECK-NEXT:    [[SA7:%.*]] = sext i8 [[A7]] to i32
+; CHECK-NEXT:    [[SB0:%.*]] = sext i8 [[B0]] to i32
+; CHECK-NEXT:    [[SB1:%.*]] = sext i8 [[B1]] to i32
+; CHECK-NEXT:    [[SB2:%.*]] = sext i8 [[B2]] to i32
+; CHECK-NEXT:    [[SB3:%.*]] = sext i8 [[B3]] to i32
+; CHECK-NEXT:    [[SB4:%.*]] = sext i8 [[B4]] to i32
+; CHECK-NEXT:    [[SB5:%.*]] = sext i8 [[B5]] to i32
+; CHECK-NEXT:    [[SB6:%.*]] = sext i8 [[B6]] to i32
+; CHECK-NEXT:    [[SB7:%.*]] = sext i8 [[B7]] to i32
+; CHECK-NEXT:    [[M0:%.*]] = mul nsw i32 [[SA0]], [[SB0]]
+; CHECK-NEXT:    [[SUM0:%.*]] = add nsw i32 [[ACC:%.*]], [[M0]]
+; CHECK-NEXT:    [[M1:%.*]] = mul nsw i32 [[SA1]], [[SB1]]
+; CHECK-NEXT:    [[SUM1:%.*]] = add nsw i32 [[SUM0]], [[M1]]
+; CHECK-NEXT:    [[M2:%.*]] = mul nsw i32 [[SA2]], [[SB2]]
+; CHECK-NEXT:    [[SUM2:%.*]] = add nsw i32 [[SUM1]], [[M2]]
+; CHECK-NEXT:    [[M3:%.*]] = mul nsw i32 [[SA3]], [[SB3]]
+; CHECK-NEXT:    [[SUM3:%.*]] = add nsw i32 [[SUM2]], [[M3]]
+; CHECK-NEXT:    [[M4:%.*]] = mul nsw i32 [[SA4]], [[SB4]]
+; CHECK-NEXT:    [[SUM4:%.*]] = add nsw i32 [[SUM3]], [[M4]]
+; CHECK-NEXT:    [[M5:%.*]] = mul nsw i32 [[SA5]], [[SB5]]
+; CHECK-NEXT:    [[SUM5:%.*]] = add nsw i32 [[SUM4]], [[M5]]
+; CHECK-NEXT:    [[M6:%.*]] = mul nsw i32 [[SA6]], [[SB6]]
+; CHECK-NEXT:    [[SUM6:%.*]] = add nsw i32 [[SUM5]], [[M6]]
+; CHECK-NEXT:    [[M7:%.*]] = mul nsw i32 [[SA7]], [[SB7]]
+; CHECK-NEXT:    [[SUM7:%.*]] = add nsw i32 [[SUM6]], [[M7]]
+; CHECK-NEXT:    ret i32 [[SUM7]]
+entry:
+  %va = load <8 x i8>, ptr addrspace(3) %a, align 8
+  %vb = load <8 x i8>, ptr addrspace(3) %b, align 8
+  %a0 = extractelement <8 x i8> %va, i32 0
+  %a1 = extractelement <8 x i8> %va, i32 1
+  %a2 = extractelement <8 x i8> %va, i32 2
+  %a3 = extractelement <8 x i8> %va, i32 3
+  %a4 = extractelement <8 x i8> %va, i32 4
+  %a5 = extractelement <8 x i8> %va, i32 5
+  %a6 = extractelement <8 x i8> %va, i32 6
+  %a7 = extractelement <8 x i8> %va, i32 7
+  %b0 = extractelement <8 x i8> %vb, i32 0
+  %b1 = extractelement <8 x i8> %vb, i32 1
+  %b2 = extractelement <8 x i8> %vb, i32 2
+  %b3 = extractelement <8 x i8> %vb, i32 3
+  %b4 = extractelement <8 x i8> %vb, i32 4
+  %b5 = extractelement <8 x i8> %vb, i32 5
+  %b6 = extractelement <8 x i8> %vb, i32 6
+  %b7 = extractelement <8 x i8> %vb, i32 7
+  %sa0 = sext i8 %a0 to i32
+  %sa1 = sext i8 %a1 to i32
+  %sa2 = sext i8 %a2 to i32
+  %sa3 = sext i8 %a3 to i32
+  %sa4 = sext i8 %a4 to i32
+  %sa5 = sext i8 %a5 to i32
+  %sa6 = sext i8 %a6 to i32
+  %sa7 = sext i8 %a7 to i32
+  %sb0 = sext i8 %b0 to i32
+  %sb1 = sext i8 %b1 to i32
+  %sb2 = sext i8 %b2 to i32
+  %sb3 = sext i8 %b3 to i32
+  %sb4 = sext i8 %b4 to i32
+  %sb5 = sext i8 %b5 to i32
+  %sb6 = sext i8 %b6 to i32
+  %sb7 = sext i8 %b7 to i32
+  %m0 = mul nsw i32 %sa0, %sb0
+  %sum0 = add nsw i32 %acc, %m0
+  %m1 = mul nsw i32 %sa1, %sb1
+  %sum1 = add nsw i32 %sum0, %m1
+  %m2 = mul nsw i32 %sa2, %sb2
+  %sum2 = add nsw i32 %sum1, %m2
+  %m3 = mul nsw i32 %sa3, %sb3
+  %sum3 = add nsw i32 %sum2, %m3
+  %m4 = mul nsw i32 %sa4, %sb4
+  %sum4 = add nsw i32 %sum3, %m4
+  %m5 = mul nsw i32 %sa5, %sb5
+  %sum5 = add nsw i32 %sum4, %m5
+  %m6 = mul nsw i32 %sa6, %sb6
+  %sum6 = add nsw i32 %sum5, %m6
+  %m7 = mul nsw i32 %sa7, %sb7
+  %sum7 = add nsw i32 %sum6, %m7
+  ret i32 %sum7
+}
+
+; Keep the minimal scalar-load form scalar too. This catches the two-lane
+; unpack shape independently of vector-load extraction.
+define i32 @dp4a_like_i8(ptr %a, ptr %b) {
+; CHECK-LABEL: @dp4a_like_i8(
+; CHECK-NEXT:  entry:
+; CHECK-NEXT:    [[A0:%.*]] = load i8, ptr [[A:%.*]], align 4
+; CHECK-NEXT:    [[A1P:%.*]] = getelementptr i8, ptr [[A]], i64 1
+; CHECK-NEXT:    [[A1:%.*]] = load i8, ptr [[A1P]], align 1
+; CHECK-NEXT:    [[B0:%.*]] = load i8, ptr [[B:%.*]], align 4
+; CHECK-NEXT:    [[B1P:%.*]] = getelementptr i8, ptr [[B]], i64 1
+; CHECK-NEXT:    [[B1:%.*]] = load i8, ptr [[B1P]], align 1
+; CHECK-NEXT:    [[SA0:%.*]] = sext i8 [[A0]] to i32
+; CHECK-NEXT:    [[SA1:%.*]] = sext i8 [[A1]] to i32
+; CHECK-NEXT:    [[SB0:%.*]] = sext i8 [[B0]] to i32
+; CHECK-NEXT:    [[SB1:%.*]] = sext i8 [[B1]] to i32
+; CHECK-NEXT:    [[M0:%.*]] = mul i32 [[SA0]], [[SB0]]
+; CHECK-NEXT:    [[M1:%.*]] = mul i32 [[SA1]], [[SB1]]
+; CHECK-NEXT:    [[SUM:%.*]] = add i32 [[M0]], [[M1]]
+; CHECK-NEXT:    ret i32 [[SUM]]
+;
+entry:
+  %a0 = load i8, ptr %a, align 4
+  %a1p = getelementptr i8, ptr %a, i64 1
+  %a1 = load i8, ptr %a1p, align 1
+  %b0 = load i8, ptr %b, align 4
+  %b1p = getelementptr i8, ptr %b, i64 1
+  %b1 = load i8, ptr %b1p, align 1
+  %sa0 = sext i8 %a0 to i32
+  %sa1 = sext i8 %a1 to i32
+  %sb0 = sext i8 %b0 to i32
+  %sb1 = sext i8 %b1 to i32
+  %m0 = mul i32 %sa0, %sb0
+  %m1 = mul i32 %sa1, %sb1
+  %sum = add i32 %m0, %m1
+  ret i32 %sum
+}
diff --git a/llvm/test/Transforms/SLPVectorizer/NVPTX/pair-copy-shapes.ll b/llvm/test/Transforms/SLPVectorizer/NVPTX/pair-copy-shapes.ll
new file mode 100644
index 0000000000000..082083a2a667d
--- /dev/null
+++ b/llvm/test/Transforms/SLPVectorizer/NVPTX/pair-copy-shapes.ll
@@ -0,0 +1,39 @@
+; RUN: opt < %s -passes=slp-vectorizer -S -mtriple=nvptx64-nvidia-cuda -mcpu=sm_100 | FileCheck %s
+
+target triple = "nvptx64-nvidia-cuda"
+
+define void @pair8_copy(ptr addrspace(1) %in, ptr addrspace(1) %out) {
+; CHECK-LABEL: @pair8_copy(
+; CHECK-NEXT:  entry:
+; CHECK-NEXT:    [[TMP0:%.*]] = load <2 x i32>, ptr addrspace(1) [[IN:%.*]], align 8
+; CHECK-NEXT:    store <2 x i32> [[TMP0]], ptr addrspace(1) [[OUT:%.*]], align 8
+; CHECK-NEXT:    ret void
+;
+entry:
+  %x = load i32, ptr addrspace(1) %in, align 8
+  %in1 = getelementptr i8, ptr addrspace(1) %in, i64 4
+  %y = load i32, ptr addrspace(1) %in1, align 4
+  store i32 %x, ptr addrspace(1) %out, align 8
+  %out1 = getelementptr i8, ptr addrspace(1) %out, i64 4
+  store i32 %y, ptr addrspace(1) %out1, align 4
+  ret void
+}
+
+define void @pair8_transform_add_sub(ptr addrspace(1) %in, ptr addrspace(1) %out,
+                                     i32 %salt) {
+; CHECK-LABEL: @pair8_transform_add_sub(
+; CHECK-NOT: add <2 x i32>
+; CHECK-NOT: sub <2 x i32>
+; CHECK: ret void
+entry:
+  %x = load i32, ptr addrspace(1) %in, align 8
+  %in1 = getelementptr i8, ptr addrspace(1) %in, i64 4
+  %y = load i32, ptr addrspace(1) %in1, align 4
+  %bit = and i32 %salt, 1
+  %nx = add i32 %x, %bit
+  %ny = sub i32 %y, %bit
+  store i32 %nx, ptr addrspace(1) %out, align 8
+  %out1 = getelementptr i8, ptr addrspace(1) %out, i64 4
+  store i32 %ny, ptr addrspace(1) %out1, align 4
+  ret void
+}
diff --git a/llvm/test/Transforms/SLPVectorizer/NVPTX/row-overhead.ll b/llvm/test/Transforms/SLPVectorizer/NVPTX/row-overhead.ll
new file mode 100644
index 0000000000000..5e8916281deb3
--- /dev/null
+++ b/llvm/test/Transforms/SLPVectorizer/NVPTX/row-overhead.ll
@@ -0,0 +1,225 @@
+; RUN: opt < %s -passes=slp-vectorizer -S -mtriple=nvptx64-nvidia-cuda -mcpu=sm_100 | FileCheck %s
+
+; Keep the row/control values scalar across the branchy tail merge.
+
+target triple = "nvptx64-nvidia-cuda"
+
+define void @slp_row_overhead_min(i32 %nnz, ptr addrspace(1) %rows,
+                                  ptr addrspace(1) %out, i1 %full,
+                                  i1 %aligned) {
+; CHECK-LABEL: @slp_row_overhead_min(
+; CHECK-NOT: phi <{{[0-9]+}} x i32>
+; CHECK-NOT: icmp {{.*}} <{{[0-9]+}} x i32>
+; CHECK-NOT: add {{.*}} <{{[0-9]+}} x i32>
+; CHECK-NOT: sub {{.*}} <{{[0-9]+}} x i32>
+; CHECK-NOT: store <{{[0-9]+}} x i32>
+; CHECK: ret void
+entry:
+  %ctaid = tail call i32 @llvm.nvvm.read.ptx.sreg.ctaid.x()
+  %tid = tail call i32 @llvm.nvvm.read.ptx.sreg.tid.x()
+  %block.base = mul nuw nsw i32 %ctaid, 896
+  %block.base64 = zext i32 %block.base to i64
+  %row.block = getelementptr inbounds i32, ptr addrspace(1) %rows, i64 %block.base64
+  br i1 %full, label %full.block, label %tail.block
+
+full.block:
+  br i1 %aligned, label %aligned.loads, label %unaligned.loads
+
+aligned.loads:
+  %lane.base.a = mul nuw nsw i32 %tid, 7
+  %lane.base.a64 = zext i32 %lane.base.a to i64
+  %a0p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %lane.base.a64
+  %a0 = load i32, ptr addrspace(1) %a0p, align 4
+  %a1i = add nuw nsw i32 %lane.base.a, 1
+  %a1i64 = zext i32 %a1i to i64
+  %a1p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %a1i64
+  %a1 = load i32, ptr addrspace(1) %a1p, align 4
+  %a2i = add nuw nsw i32 %lane.base.a, 2
+  %a2i64 = zext i32 %a2i to i64
+  %a2p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %a2i64
+  %a2 = load i32, ptr addrspace(1) %a2p, align 4
+  %a3i = add nuw nsw i32 %lane.base.a, 3
+  %a3i64 = zext i32 %a3i to i64
+  %a3p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %a3i64
+  %a3 = load i32, ptr addrspace(1) %a3p, align 4
+  %a4i = add nuw nsw i32 %lane.base.a, 4
+  %a4i64 = zext i32 %a4i to i64
+  %a4p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %a4i64
+  %a4 = load i32, ptr addrspace(1) %a4p, align 4
+  %a5i = add nuw nsw i32 %lane.base.a, 5
+  %a5i64 = zext i32 %a5i to i64
+  %a5p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %a5i64
+  %a5 = load i32, ptr addrspace(1) %a5p, align 4
+  %a6i = add nuw nsw i32 %lane.base.a, 6
+  %a6i64 = zext i32 %a6i to i64
+  %a6p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %a6i64
+  %a6 = load i32, ptr addrspace(1) %a6p, align 4
+  br label %merge
+
+unaligned.loads:
+  %lane.base.b = mul nuw nsw i32 %tid, 7
+  %lane.base.b64 = zext i32 %lane.base.b to i64
+  %b0p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %lane.base.b64
+  %b0 = load i32, ptr addrspace(1) %b0p, align 4
+  %b1i = add nuw nsw i32 %lane.base.b, 1
+  %b1i64 = zext i32 %b1i to i64
+  %b1p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %b1i64
+  %b1 = load i32, ptr addrspace(1) %b1p, align 4
+  %b2i = add nuw nsw i32 %lane.base.b, 2
+  %b2i64 = zext i32 %b2i to i64
+  %b2p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %b2i64
+  %b2 = load i32, ptr addrspace(1) %b2p, align 4
+  %b3i = add nuw nsw i32 %lane.base.b, 3
+  %b3i64 = zext i32 %b3i to i64
+  %b3p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %b3i64
+  %b3 = load i32, ptr addrspace(1) %b3p, align 4
+  %b4i = add nuw nsw i32 %lane.base.b, 4
+  %b4i64 = zext i32 %b4i to i64
+  %b4p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %b4i64
+  %b4 = load i32, ptr addrspace(1) %b4p, align 4
+  %b5i = add nuw nsw i32 %lane.base.b, 5
+  %b5i64 = zext i32 %b5i to i64
+  %b5p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %b5i64
+  %b5 = load i32, ptr addrspace(1) %b5p, align 4
+  %b6i = add nuw nsw i32 %lane.base.b, 6
+  %b6i64 = zext i32 %b6i to i64
+  %b6p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %b6i64
+  %b6 = load i32, ptr addrspace(1) %b6p, align 4
+  br label %merge
+
+tail.block:
+  %tail.left = sub nsw i32 %nnz, %block.base
+  %tail.last.i = add nsw i32 %tail.left, -1
+  %tail.last.i64 = sext i32 %tail.last.i to i64
+  %lastp = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %tail.last.i64
+  %last = load i32, ptr addrspace(1) %lastp, align 4
+  %lane.base.t = mul nuw nsw i32 %tid, 7
+  %tail.rem = sub nsw i32 %tail.left, %lane.base.t
+  %tail.has0 = icmp sgt i32 %tail.rem, 0
+  br i1 %tail.has0, label %tail.load0, label %tail.merge0
+
+tail.load0:
+  %t0i64 = zext i32 %lane.base.t to i64
+  %t0p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %t0i64
+  %t0 = load i32, ptr addrspace(1) %t0p, align 4
+  br label %tail.merge0
+
+tail.merge0:
+  %c0 = phi i32 [ %t0, %tail.load0 ], [ %last, %tail.block ]
+  %tail.has1 = icmp sgt i32 %tail.rem, 1
+  br i1 %tail.has1, label %tail.load1, label %tail.merge1
+
+tail.load1:
+  %t1i = add nuw nsw i32 %lane.base.t, 1
+  %t1i64 = zext i32 %t1i to i64
+  %t1p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %t1i64
+  %t1 = load i32, ptr addrspace(1) %t1p, align 4
+  br label %tail.merge1
+
+tail.merge1:
+  %c1 = phi i32 [ %t1, %tail.load1 ], [ %last, %tail.merge0 ]
+  %tail.has2 = icmp sgt i32 %tail.rem, 2
+  br i1 %tail.has2, label %tail.load2, label %tail.merge2
+
+tail.load2:
+  %t2i = add nuw nsw i32 %lane.base.t, 2
+  %t2i64 = zext i32 %t2i to i64
+  %t2p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %t2i64
+  %t2 = load i32, ptr addrspace(1) %t2p, align 4
+  br label %tail.merge2
+
+tail.merge2:
+  %c2 = phi i32 [ %t2, %tail.load2 ], [ %last, %tail.merge1 ]
+  %tail.has3 = icmp sgt i32 %tail.rem, 3
+  br i1 %tail.has3, label %tail.load3, label %tail.merge3
+
+tail.load3:
+  %t3i = add nuw nsw i32 %lane.base.t, 3
+  %t3i64 = zext i32 %t3i to i64
+  %t3p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %t3i64
+  %t3 = load i32, ptr addrspace(1) %t3p, align 4
+  br label %tail.merge3
+
+tail.merge3:
+  %c3 = phi i32 [ %t3, %tail.load3 ], [ %last, %tail.merge2 ]
+  %tail.has4 = icmp sgt i32 %tail.rem, 4
+  br i1 %tail.has4, label %tail.load4, label %tail.merge4
+
+tail.load4:
+  %t4i = add nuw nsw i32 %lane.base.t, 4
+  %t4i64 = zext i32 %t4i to i64
+  %t4p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %t4i64
+  %t4 = load i32, ptr addrspace(1) %t4p, align 4
+  br label %tail.merge4
+
+tail.merge4:
+  %c4 = phi i32 [ %t4, %tail.load4 ], [ %last, %tail.merge3 ]
+  %tail.has5 = icmp sgt i32 %tail.rem, 5
+  br i1 %tail.has5, label %tail.load5, label %tail.merge5
+
+tail.load5:
+  %t5i = add nuw nsw i32 %lane.base.t, 5
+  %t5i64 = zext i32 %t5i to i64
+  %t5p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %t5i64
+  %t5 = load i32, ptr addrspace(1) %t5p, align 4
+  br label %tail.merge5
+
+tail.merge5:
+  %c5 = phi i32 [ %t5, %tail.load5 ], [ %last, %tail.merge4 ]
+  %tail.has6 = icmp sgt i32 %tail.rem, 6
+  br i1 %tail.has6, label %tail.load6, label %merge
+
+tail.load6:
+  %t6i = add nuw nsw i32 %lane.base.t, 6
+  %t6i64 = zext i32 %t6i to i64
+  %t6p = getelementptr inbounds i32, ptr addrspace(1) %row.block, i64 %t6i64
+  %t6 = load i32, ptr addrspace(1) %t6p, align 4
+  br label %merge
+
+merge:
+  %r6 = phi i32 [ %a6, %aligned.loads ], [ %b6, %unaligned.loads ], [ %t6, %tail.load6 ], [ %last, %tail.merge5 ]
+  %r5 = phi i32 [ %a5, %aligned.loads ], [ %b5, %unaligned.loads ], [ %c5, %tail.load6 ], [ %c5, %tail.merge5 ]
+  %r4 = phi i32 [ %a4, %aligned.loads ], [ %b4, %unaligned.loads ], [ %c4, %tail.load6 ], [ %c4, %tail.merge5 ]
+  %r3 = phi i32 [ %a3, %aligned.loads ], [ %b3, %unaligned.loads ], [ %c3, %tail.load6 ], [ %c3, %tail.merge5 ]
+  %r2 = phi i32 [ %a2, %aligned.loads ], [ %b2, %unaligned.loads ], [ %c2, %tail.load6 ], [ %c2, %tail.merge5 ]
+  %r1 = phi i32 [ %a1, %aligned.loads ], [ %b1, %unaligned.loads ], [ %c1, %tail.load6 ], [ %c1, %tail.merge5 ]
+  %r0 = phi i32 [ %a0, %aligned.loads ], [ %b0, %unaligned.loads ], [ %c0, %tail.load6 ], [ %c0, %tail.merge5 ]
+  %sh0 = tail call i32 asm sideeffect "shfl.sync.idx.b32 $0, $1, $2, 31, -1;", "=r,r,r"(i32 %r0, i32 0)
+  %sh6 = tail call i32 asm sideeffect "shfl.sync.idx.b32 $0, $1, $2, 31, -1;", "=r,r,r"(i32 %r6, i32 31)
+  %cmp.ends = icmp eq i32 %sh0, %sh6
+  br i1 %cmp.ends, label %same.row, label %split.row
+
+same.row:
+  store i32 %sh0, ptr addrspace(1) %out, align 4
+  ret void
+
+split.row:
+  %up6 = tail call i32 asm sideeffect "shfl.sync.up.b32 $0, $1, $2, 0, -1;", "=r,r,r"(i32 %r6, i32 1)
+  %d0 = sub nsw i32 %r0, %sh0
+  %d1 = sub nsw i32 %r1, %sh0
+  %d2 = sub nsw i32 %r2, %sh0
+  %d3 = sub nsw i32 %r3, %sh0
+  %d4 = sub nsw i32 %r4, %sh0
+  %d5 = sub nsw i32 %r5, %sh0
+  %cmp06 = icmp ne i32 %r0, %r6
+  %cmpup = icmp ne i32 %up6, %r0
+  %cmp.int = zext i1 %cmp06 to i32
+  %cmpup.int = zext i1 %cmpup to i32
+  %mix0 = add i32 %d0, %cmp.int
+  %mix1 = add i32 %d5, %cmpup.int
+  store i32 %mix0, ptr addrspace(1) %out, align 4
+  %out1 = getelementptr inbounds i32, ptr addrspace(1) %out, i64 1
+  store i32 %d1, ptr addrspace(1) %out1, align 4
+  %out2 = getelementptr inbounds i32, ptr addrspace(1) %out, i64 2
+  store i32 %d2, ptr addrspace(1) %out2, align 4
+  %out3 = getelementptr inbounds i32, ptr addrspace(1) %out, i64 3
+  store i32 %d3, ptr addrspace(1) %out3, align 4
+  %out4 = getelementptr inbounds i32, ptr addrspace(1) %out, i64 4
+  store i32 %d4, ptr addrspace(1) %out4, align 4
+  %out5 = getelementptr inbounds i32, ptr addrspace(1) %out, i64 5
+  store i32 %mix1, ptr addrspace(1) %out5, align 4
+  ret void
+}
+
+declare i32 @llvm.nvvm.read.ptx.sreg.ctaid.x()
+declare i32 @llvm.nvvm.read.ptx.sreg.tid.x()
diff --git a/llvm/test/Transforms/SLPVectorizer/NVPTX/v2i16-scalar-uses.ll b/llvm/test/Transforms/SLPVectorizer/NVPTX/v2i16-scalar-uses.ll
new file mode 100644
index 0000000000000..8c5f0a78bc0c5
--- /dev/null
+++ b/llvm/test/Transforms/SLPVectorizer/NVPTX/v2i16-scalar-uses.ll
@@ -0,0 +1,79 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py
+; RUN: opt < %s -passes=slp-vectorizer -S -mtriple=nvptx64-nvidia-cuda -mcpu=sm_100 | FileCheck %s
+
+target triple = "nvptx64-nvidia-cuda"
+
+declare i32 @llvm.nvvm.shfl.idx.i32(i32, i32, i32)
+
+; Packed i16 values used by scalar intrinsics and scalar integer stores should
+; not be vectorized just because v2i16 is a cheap packed memory type.
+define void @i16_pair_intrinsic_uses(ptr %out, ptr %side, ptr %a, ptr %b,
+                                     i16 %pred) {
+; CHECK-LABEL: @i16_pair_intrinsic_uses(
+; CHECK-NEXT:  entry:
+; CHECK-NEXT:    [[COND:%.*]] = and i16 [[PRED:%.*]], 1
+; CHECK-NEXT:    [[TOBOOL:%.*]] = icmp eq i16 [[COND]], 0
+; CHECK-NEXT:    br i1 [[TOBOOL]], label [[ELSE:%.*]], label [[THEN:%.*]]
+; CHECK:       then:
+; CHECK-NEXT:    [[A0:%.*]] = load i16, ptr [[A:%.*]], align 4
+; CHECK-NEXT:    [[A1P:%.*]] = getelementptr i16, ptr [[A]], i64 1
+; CHECK-NEXT:    [[A1:%.*]] = load i16, ptr [[A1P]], align 2
+; CHECK-NEXT:    br label [[JOIN:%.*]]
+; CHECK:       else:
+; CHECK-NEXT:    [[B0:%.*]] = load i16, ptr [[B:%.*]], align 4
+; CHECK-NEXT:    [[B1P:%.*]] = getelementptr i16, ptr [[B]], i64 1
+; CHECK-NEXT:    [[B1:%.*]] = load i16, ptr [[B1P]], align 2
+; CHECK-NEXT:    br label [[JOIN]]
+; CHECK:       join:
+; CHECK-NEXT:    [[X0:%.*]] = phi i16 [ [[A0]], [[THEN]] ], [ [[B0]], [[ELSE]] ]
+; CHECK-NEXT:    [[X1:%.*]] = phi i16 [ [[A1]], [[THEN]] ], [ [[B1]], [[ELSE]] ]
+; CHECK-NEXT:    store i16 [[X0]], ptr [[OUT:%.*]], align 4
+; CHECK-NEXT:    [[OUT1:%.*]] = getelementptr i16, ptr [[OUT]], i64 1
+; CHECK-NEXT:    store i16 [[X1]], ptr [[OUT1]], align 2
+; CHECK-NEXT:    [[E0:%.*]] = sext i16 [[X0]] to i32
+; CHECK-NEXT:    [[E1:%.*]] = sext i16 [[X1]] to i32
+; CHECK-NEXT:    [[P32:%.*]] = sext i16 [[PRED]] to i32
+; CHECK-NEXT:    [[Y0:%.*]] = call i32 @llvm.nvvm.shfl.idx.i32(i32 [[E0]], i32 [[P32]], i32 31)
+; CHECK-NEXT:    [[Y1:%.*]] = call i32 @llvm.nvvm.shfl.idx.i32(i32 [[E1]], i32 [[P32]], i32 31)
+; CHECK-NEXT:    [[S0:%.*]] = add i32 [[Y0]], [[E1]]
+; CHECK-NEXT:    [[S1:%.*]] = add i32 [[Y1]], [[E0]]
+; CHECK-NEXT:    store i32 [[S0]], ptr [[SIDE:%.*]], align 8
+; CHECK-NEXT:    [[SIDE1:%.*]] = getelementptr i32, ptr [[SIDE]], i64 1
+; CHECK-NEXT:    store i32 [[S1]], ptr [[SIDE1]], align 4
+; CHECK-NEXT:    ret void
+;
+entry:
+  %cond = and i16 %pred, 1
+  %tobool = icmp eq i16 %cond, 0
+  br i1 %tobool, label %else, label %then
+
+then:
+  %a0 = load i16, ptr %a, align 4
+  %a1p = getelementptr i16, ptr %a, i64 1
+  %a1 = load i16, ptr %a1p, align 2
+  br label %join
+
+else:
+  %b0 = load i16, ptr %b, align 4
+  %b1p = getelementptr i16, ptr %b, i64 1
+  %b1 = load i16, ptr %b1p, align 2
+  br label %join
+
+join:
+  %x0 = phi i16 [ %a0, %then ], [ %b0, %else ]
+  %x1 = phi i16 [ %a1, %then ], [ %b1, %else ]
+  store i16 %x0, ptr %out, align 4
+  %out1 = getelementptr i16, ptr %out, i64 1
+  store i16 %x1, ptr %out1, align 2
+  %e0 = sext i16 %x0 to i32
+  %e1 = sext i16 %x1 to i32
+  %p32 = sext i16 %pred to i32
+  %y0 = call i32 @llvm.nvvm.shfl.idx.i32(i32 %e0, i32 %p32, i32 31)
+  %y1 = call i32 @llvm.nvvm.shfl.idx.i32(i32 %e1, i32 %p32, i32 31)
+  %s0 = add i32 %y0, %e1
+  %s1 = add i32 %y1, %e0
+  store i32 %s0, ptr %side, align 8
+  %side1 = getelementptr i32, ptr %side, i64 1
+  store i32 %s1, ptr %side1, align 4
+  ret void
+}

>From 11248bc04c4d3b839c2b4dde2400452ccaeb295c Mon Sep 17 00:00:00 2001
From: Daniel Donenfeld <ddonenfeld at nvidia.com>
Date: Tue, 15 Sep 2026 15:46:12 +0000
Subject: [PATCH 2/2] Fix issues

---
 .../Target/NVPTX/NVPTXTargetTransformInfo.cpp | 42 ++++++++++---------
 .../Target/NVPTX/NVPTXTargetTransformInfo.h   |  3 +-
 .../NVPTX/ordered-reduction-fma-fusion.ll     | 28 +++++++++----
 3 files changed, 44 insertions(+), 29 deletions(-)

diff --git a/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.cpp b/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.cpp
index bf0c43f32cea7..6798f92c150d2 100644
--- a/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.cpp
@@ -538,8 +538,9 @@ static InstructionCost getVectorRegisterPieceCost(Type *Ty,
   return NumPieces;
 }
 
-static InstructionCost getNonNativeVectorOpPenalty(
-    Type *Ty, const DataLayout &DL, const TargetLoweringBase *TLI) {
+static InstructionCost
+getNonNativeVectorOpPenalty(Type *Ty, const DataLayout &DL,
+                            const TargetLoweringBase *TLI) {
   auto *VTy = dyn_cast<FixedVectorType>(Ty);
   if (!VTy || isCheapPTXVectorInsertExtract(Ty, DL, TLI))
     return 0;
@@ -549,9 +550,10 @@ static InstructionCost getNonNativeVectorOpPenalty(
   return VTy->getNumElements() + getVectorRegisterPieceCost(Ty, DL);
 }
 
-static InstructionCost getNonNativeVectorShufflePenalty(
-    VectorType *DstTy, VectorType *SrcTy, VectorType *SubTp,
-    const DataLayout &DL, const TargetLoweringBase *TLI) {
+static InstructionCost
+getNonNativeVectorShufflePenalty(VectorType *DstTy, VectorType *SrcTy,
+                                 VectorType *SubTp, const DataLayout &DL,
+                                 const TargetLoweringBase *TLI) {
   InstructionCost Cost = 0;
   for (VectorType *Ty : {DstTy, SrcTy, SubTp}) {
     if (!Ty)
@@ -561,14 +563,14 @@ static InstructionCost getNonNativeVectorShufflePenalty(
   return Cost;
 }
 
-InstructionCost NVPTXTTIImpl::getShuffleCost(
-    TTI::ShuffleKind Kind, VectorType *DstTy, VectorType *SrcTy,
-    ArrayRef<int> Mask, TTI::TargetCostKind CostKind, int Index,
-    VectorType *SubTp, ArrayRef<const Value *> Args,
-    const Instruction *CxtI) const {
-  InstructionCost Cost =
-      BaseT::getShuffleCost(Kind, DstTy, SrcTy, Mask, CostKind, Index, SubTp,
-                            Args, CxtI);
+InstructionCost
+NVPTXTTIImpl::getShuffleCost(TTI::ShuffleKind Kind, VectorType *DstTy,
+                             VectorType *SrcTy, TTI::TargetCostKind CostKind,
+                             ArrayRef<int> Mask, int Index, VectorType *SubTp,
+                             ArrayRef<const Value *> Args,
+                             const Instruction *CxtI) const {
+  InstructionCost Cost = BaseT::getShuffleCost(Kind, DstTy, SrcTy, CostKind,
+                                               Mask, Index, SubTp, Args, CxtI);
 
   if (CostKind != TTI::TCK_RecipThroughput)
     return Cost;
@@ -695,9 +697,11 @@ InstructionCost NVPTXTTIImpl::getScalarizationOverhead(
   return Cost + BaseCost;
 }
 
-InstructionCost NVPTXTTIImpl::getCastInstrCost(
-    unsigned Opcode, Type *Dst, Type *Src, TTI::CastContextHint CCH,
-    TTI::TargetCostKind CostKind, const Instruction *I) const {
+InstructionCost NVPTXTTIImpl::getCastInstrCost(unsigned Opcode, Type *Dst,
+                                               Type *Src,
+                                               TTI::CastContextHint CCH,
+                                               TTI::TargetCostKind CostKind,
+                                               const Instruction *I) const {
   InstructionCost Cost =
       BaseT::getCastInstrCost(Opcode, Dst, Src, CCH, CostKind, I);
   if (CostKind != TTI::TCK_RecipThroughput)
@@ -713,9 +717,9 @@ InstructionCost NVPTXTTIImpl::getCastInstrCost(
         DstVTy->getNumElements() *
         BaseT::getCastInstrCost(Opcode, DstVTy->getElementType(),
                                 SrcVTy->getElementType(), CCH, CostKind, I);
-    InstructionCost LegalizedCost =
-        ScalarCost + getNonNativeVectorOpPenalty(Dst, DL, TLI) +
-        getNonNativeVectorOpPenalty(Src, DL, TLI);
+    InstructionCost LegalizedCost = ScalarCost +
+                                    getNonNativeVectorOpPenalty(Dst, DL, TLI) +
+                                    getNonNativeVectorOpPenalty(Src, DL, TLI);
     if (LegalizedCost > Cost)
       Cost = LegalizedCost;
   } else if (Opcode == Instruction::BitCast) {
diff --git a/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h b/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h
index 6e586581b1384..ecd514f02858c 100644
--- a/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h
+++ b/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h
@@ -129,8 +129,7 @@ class NVPTXTTIImpl final : public BasicTTIImplBase<NVPTXTTIImpl> {
 
   InstructionCost
   getShuffleCost(TTI::ShuffleKind Kind, VectorType *DstTy, VectorType *SrcTy,
-                 ArrayRef<int> Mask = {},
-                 TTI::TargetCostKind CostKind = TTI::TCK_RecipThroughput,
+                 TTI::TargetCostKind CostKind, ArrayRef<int> Mask = {},
                  int Index = 0, VectorType *SubTp = nullptr,
                  ArrayRef<const Value *> Args = {},
                  const Instruction *CxtI = nullptr) const override;
diff --git a/llvm/test/Transforms/SLPVectorizer/NVPTX/ordered-reduction-fma-fusion.ll b/llvm/test/Transforms/SLPVectorizer/NVPTX/ordered-reduction-fma-fusion.ll
index a553fe126ee17..be0a05c05d538 100644
--- a/llvm/test/Transforms/SLPVectorizer/NVPTX/ordered-reduction-fma-fusion.ll
+++ b/llvm/test/Transforms/SLPVectorizer/NVPTX/ordered-reduction-fma-fusion.ll
@@ -9,10 +9,14 @@
 
 define float @dot_contract(float %x) {
 ; CHECK-LABEL: @dot_contract(
-; CHECK-NEXT:    [[TMP1:%.*]] = insertelement <4 x float> poison, float [[X:%.*]], i64 0
-; CHECK-NEXT:    [[TMP2:%.*]] = shufflevector <4 x float> [[TMP1]], <4 x float> poison, <4 x i32> zeroinitializer
-; CHECK-NEXT:    [[TMP3:%.*]] = fmul contract <4 x float> <float 7.000000e+00, float 3.000000e+00, float 5.000000e+00, float 9.000000e+00>, [[TMP2]]
-; CHECK-NEXT:    [[TMP4:%.*]] = call contract float @llvm.vector.reduce.fadd.v4f32(float [[X]], <4 x float> [[TMP3]])
+; CHECK-NEXT:    [[M0:%.*]] = fmul contract float 7.000000e+00, [[X:%.*]]
+; CHECK-NEXT:    [[A0:%.*]] = fadd contract float [[M0]], [[X]]
+; CHECK-NEXT:    [[M1:%.*]] = fmul contract float 3.000000e+00, [[X]]
+; CHECK-NEXT:    [[A1:%.*]] = fadd contract float [[M1]], [[A0]]
+; CHECK-NEXT:    [[M2:%.*]] = fmul contract float 5.000000e+00, [[X]]
+; CHECK-NEXT:    [[A2:%.*]] = fadd contract float [[M2]], [[A1]]
+; CHECK-NEXT:    [[M3:%.*]] = fmul contract float 9.000000e+00, [[X]]
+; CHECK-NEXT:    [[TMP4:%.*]] = fadd contract float [[M3]], [[A2]]
 ; CHECK-NEXT:    ret float [[TMP4]]
 ;
   %m0 = fmul contract float 7.000000e+00, %x
@@ -28,10 +32,18 @@ define float @dot_contract(float %x) {
 
 define float @dot_no_contract(float %x) {
 ; CHECK-LABEL: @dot_no_contract(
-; CHECK-NEXT:    [[TMP1:%.*]] = insertelement <4 x float> poison, float [[X:%.*]], i64 0
-; CHECK-NEXT:    [[TMP2:%.*]] = shufflevector <4 x float> [[TMP1]], <4 x float> poison, <4 x i32> zeroinitializer
-; CHECK-NEXT:    [[TMP3:%.*]] = fmul <4 x float> <float 7.000000e+00, float 3.000000e+00, float 5.000000e+00, float 9.000000e+00>, [[TMP2]]
-; CHECK-NEXT:    [[TMP4:%.*]] = call float @llvm.vector.reduce.fadd.v4f32(float [[X]], <4 x float> [[TMP3]])
+; CHECK-NEXT:    [[TMP1:%.*]] = insertelement <2 x float> poison, float [[X:%.*]], i64 0
+; CHECK-NEXT:    [[TMP2:%.*]] = shufflevector <2 x float> [[TMP1]], <2 x float> poison, <2 x i32> zeroinitializer
+; CHECK-NEXT:    [[TMP3:%.*]] = fmul <2 x float> <float 3.000000e+00, float 7.000000e+00>, [[TMP2]]
+; CHECK-NEXT:    [[TMP9:%.*]] = extractelement <2 x float> [[TMP3]], i64 1
+; CHECK-NEXT:    [[A0:%.*]] = fadd float [[TMP9]], [[X]]
+; CHECK-NEXT:    [[TMP5:%.*]] = extractelement <2 x float> [[TMP3]], i64 0
+; CHECK-NEXT:    [[A1:%.*]] = fadd float [[TMP5]], [[A0]]
+; CHECK-NEXT:    [[TMP6:%.*]] = fmul <2 x float> <float 9.000000e+00, float 5.000000e+00>, [[TMP2]]
+; CHECK-NEXT:    [[TMP7:%.*]] = extractelement <2 x float> [[TMP6]], i64 1
+; CHECK-NEXT:    [[A2:%.*]] = fadd float [[TMP7]], [[A1]]
+; CHECK-NEXT:    [[TMP8:%.*]] = extractelement <2 x float> [[TMP6]], i64 0
+; CHECK-NEXT:    [[TMP4:%.*]] = fadd float [[TMP8]], [[A2]]
 ; CHECK-NEXT:    ret float [[TMP4]]
 ;
   %m0 = fmul float 7.000000e+00, %x



More information about the llvm-commits mailing list