[llvm] [AArch64][SelectionDAG] Fold SVE interleaves of splat vectors (PR #219146)

Serval MARTINOT-LAGARDE via llvm-commits llvm-commits at lists.llvm.org
Thu Aug 27 05:56:09 PDT 2026


https://github.com/Serval6 updated https://github.com/llvm/llvm-project/pull/219146

>From 4755fcf7eaa88ec373a4ca8555acee229b4e8016 Mon Sep 17 00:00:00 2001
From: Serval Martinot-Lagarde <serval.martinot-lagarde at sipearl.com>
Date: Thu, 27 Aug 2026 10:35:18 +0200
Subject: [PATCH] [AArch64][SDag] Fold SVE interleaves of splat vectors

---
 .../Target/AArch64/AArch64ISelLowering.cpp    | 109 +++++++++++++++++-
 .../AArch64/sve-interleave-of-splat.ll        |  64 ++++++++++
 2 files changed, 172 insertions(+), 1 deletion(-)
 create mode 100644 llvm/test/CodeGen/AArch64/sve-interleave-of-splat.ll

diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index e5c419c42fd2c..6e29fe8f148fe 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -22740,6 +22740,106 @@ performExtractVectorEltCombine(SDNode *N, TargetLowering::DAGCombinerInfo &DCI,
   return SDValue();
 }
 
+static SDValue getPTrueForHalfElementCount(EVT ResultVT, const SDLoc &DL,
+                                           SelectionDAG &DAG) {
+  MVT PredicateVT = MVT::INVALID_SIMPLE_VALUE_TYPE;
+  switch (ResultVT.getVectorMinNumElements()) {
+  case 2:
+    PredicateVT = MVT::nxv1i1;
+    break;
+  case 4:
+    PredicateVT = MVT::nxv2i1;
+    break;
+  case 8:
+    PredicateVT = MVT::nxv4i1;
+    break;
+  case 16:
+    PredicateVT = MVT::nxv8i1;
+    break;
+  case 32:
+    PredicateVT = MVT::nxv16i1;
+    break;
+  case 64:
+    PredicateVT = MVT::nxv32i1;
+    break;
+  case 128:
+    PredicateVT = MVT::nxv64i1;
+    break;
+  default:
+    return SDValue();
+  }
+
+  return getSVEPredicateBitCast(
+      ResultVT, getPTrue(DAG, DL, PredicateVT, AArch64SVEPredPattern::all),
+      DAG);
+}
+
+static SDValue simplifyAlternatingMask(SDNode *N, SelectionDAG &DAG) {
+  // Try to match:
+  //    t2: nxv2i1 = splat_vector Constant:i1<0>
+  //    t7: nxv2i1 = splat_vector Constant:i1<-1>
+  //  t8: nxv2i1,nxv2i1 = vector_interleave t2, t7
+  //    t0: ch,glue = EntryToken
+  //    t9: nxv4i1 = concat_vectors t8, t8:1
+
+  // Match the concatenation of an interleave of two vectors.
+  EVT VT = N->getValueType(0);
+  SDValue Interleave = N->getOperand(0);
+  if (!VT.isScalableVector() || VT.getVectorElementType() != MVT::i1 ||
+      Interleave.getNode() != N->getOperand(1).getNode() ||
+      Interleave.getResNo() != 0 || N->getOperand(1).getResNo() != 1 ||
+      Interleave.getOpcode() != ISD::VECTOR_INTERLEAVE)
+    return SDValue();
+
+  // Match an interleave of all-zero and all-one predicate splats.
+  SDValue LHS = Interleave.getOperand(0);
+  SDValue RHS = Interleave.getOperand(1);
+  ConstantSDNode *C0 = nullptr, *C1 = nullptr;
+  if (LHS.getOpcode() != ISD::SPLAT_VECTOR ||
+      RHS.getOpcode() != ISD::SPLAT_VECTOR ||
+      !(C0 = dyn_cast<ConstantSDNode>(LHS.getOperand(0))) ||
+      !(C1 = dyn_cast<ConstantSDNode>(RHS.getOperand(0))) ||
+      !((C0->isZero() && C1->isAllOnes()) || (C0->isAllOnes() && C1->isZero())))
+    return SDValue();
+
+  // Materialize the alternating predicate directly using PTRUE (or its
+  // complement). Use it with a 2x wider element in order to effectively have
+  // an alternating mask.
+  SDValue PTrue = getPTrueForHalfElementCount(VT, N, DAG);
+  if (!PTrue)
+    return SDValue();
+
+  return C0->isZero() ? DAG.getNOT(N, PTrue, VT) : PTrue;
+}
+
+static SDValue simplifyAlternatingSplat(SDNode *N, SelectionDAG &DAG) {
+  EVT VT = N->getValueType(0);
+  SDValue Interleave = N->getOperand(0);
+  if (!VT.isScalableVT() || Interleave.getOpcode() != ISD::VECTOR_INTERLEAVE ||
+      Interleave.getNode() != N->getOperand(1).getNode() ||
+      Interleave.getResNo() != 0 || N->getOperand(1).getResNo() != 1 ||
+      Interleave.getOperand(0).getOpcode() != ISD::SPLAT_VECTOR ||
+      Interleave.getOperand(1).getOpcode() != ISD::SPLAT_VECTOR)
+    return SDValue();
+
+  const TargetLowering &TLI = DAG.getTargetLoweringInfo();
+  EVT ScalarTy = VT.getVectorElementType();
+  if (!TLI.isTypeLegal(VT) || ScalarTy == MVT::i8 || ScalarTy == MVT::i16)
+    return SDValue();
+
+  SDValue ActiveVal = Interleave.getOperand(0).getOperand(0);
+  SDValue InactiveVal = Interleave.getOperand(1).getOperand(0);
+  SDValue PTrue = getPTrueForHalfElementCount(
+      MVT::getScalableVectorVT(MVT::i1, VT.getVectorMinNumElements()), N, DAG);
+  if (!PTrue)
+    return SDValue();
+
+  SDLoc DL(N);
+  SDValue InactiveVector = DAG.getSplatVector(VT, DL, InactiveVal);
+  return DAG.getNode(AArch64ISD::DUP_MERGE_PASSTHRU, DL, VT, PTrue, ActiveVal,
+                     InactiveVector);
+}
+
 static SDValue performConcatVectorsCombine(SDNode *N,
                                            TargetLowering::DAGCombinerInfo &DCI,
                                            SelectionDAG &DAG) {
@@ -22764,8 +22864,15 @@ static SDValue performConcatVectorsCombine(SDNode *N,
     return DAG.getNode(AArch64ISD::TRN1, DL, VT, Op0MoreElems, Op1MoreElems);
   }
 
-  if (VT.isScalableVector())
+  if (VT.isScalableVector()) {
+    if (SDValue V = simplifyAlternatingMask(N, DAG))
+      return V;
+
+    if (SDValue V = simplifyAlternatingSplat(N, DAG))
+      return V;
+
     return SDValue();
+  }
 
   if (N->getNumOperands() == 2 && N0Opc == ISD::TRUNCATE &&
       N1Opc == ISD::TRUNCATE) {
diff --git a/llvm/test/CodeGen/AArch64/sve-interleave-of-splat.ll b/llvm/test/CodeGen/AArch64/sve-interleave-of-splat.ll
new file mode 100644
index 0000000000000..1a1d3e87a0b65
--- /dev/null
+++ b/llvm/test/CodeGen/AArch64/sve-interleave-of-splat.ll
@@ -0,0 +1,64 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 5
+; RUN: llc < %s | FileCheck %s
+
+target datalayout = "e-m:e-i8:8:32-i16:16:32-i64:64-i128:128-n32:64-S128-Fn32"
+target triple = "aarch64"
+
+define <vscale x 2 x float> @interleave2_nxv2f32(float %a, float %b) #0 {
+; CHECK-LABEL: interleave2_nxv2f32:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    // kill: def $s1 killed $s1 def $z1
+; CHECK-NEXT:    ptrue p0.d
+; CHECK-NEXT:    mov z1.s, s1
+; CHECK-NEXT:    punpklo p0.h, p0.b
+; CHECK-NEXT:    mov z1.s, p0/m, s0
+; CHECK-NEXT:    mov z0.d, z1.d
+; CHECK-NEXT:    ret
+  %a.insert = insertelement <vscale x 1 x float> poison, float %a, i64 0
+  %a.splat = shufflevector <vscale x 1 x float> %a.insert, <vscale x 1 x float> poison, <vscale x 1 x i32> zeroinitializer
+  %b.insert = insertelement <vscale x 1 x float> poison, float %b, i64 0
+  %b.splat = shufflevector <vscale x 1 x float> %b.insert, <vscale x 1 x float> poison, <vscale x 1 x i32> zeroinitializer
+  %res = call <vscale x 2 x float> @llvm.vector.interleave2.nxv2f32(<vscale x 1 x float> %a.splat, <vscale x 1 x float> %b.splat)
+  ret <vscale x 2 x float> %res
+}
+
+define <vscale x 4 x float> @interleave2_nxv4f32(float %a, float %b) #0 {
+; CHECK-LABEL: interleave2_nxv4f32:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    // kill: def $s1 killed $s1 def $z1
+; CHECK-NEXT:    ptrue p0.d
+; CHECK-NEXT:    mov z1.s, s1
+; CHECK-NEXT:    mov z1.s, p0/m, s0
+; CHECK-NEXT:    mov z0.d, z1.d
+; CHECK-NEXT:    ret
+  %a.insert = insertelement <vscale x 2 x float> poison, float %a, i64 0
+  %a.splat = shufflevector <vscale x 2 x float> %a.insert, <vscale x 2 x float> poison, <vscale x 2 x i32> zeroinitializer
+  %b.insert = insertelement <vscale x 2 x float> poison, float %b, i64 0
+  %b.splat = shufflevector <vscale x 2 x float> %b.insert, <vscale x 2 x float> poison, <vscale x 2 x i32> zeroinitializer
+  %res = call <vscale x 4 x float> @llvm.vector.interleave2.nxv4f32(<vscale x 2 x float> %a.splat, <vscale x 2 x float> %b.splat)
+  ret <vscale x 4 x float> %res
+}
+
+define <vscale x 8 x float> @interleave2_nxv8f32(float %a, float %b) #0 {
+; CHECK-LABEL: interleave2_nxv8f32:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    // kill: def $s1 killed $s1 def $z1
+; CHECK-NEXT:    // kill: def $s0 killed $s0 def $z0
+; CHECK-NEXT:    mov z2.s, s0
+; CHECK-NEXT:    mov z1.s, s1
+; CHECK-NEXT:    zip1 z0.s, z2.s, z1.s
+; CHECK-NEXT:    zip2 z1.s, z2.s, z1.s
+; CHECK-NEXT:    ret
+  %a.insert = insertelement <vscale x 4 x float> poison, float %a, i64 0
+  %a.splat = shufflevector <vscale x 4 x float> %a.insert, <vscale x 4 x float> poison, <vscale x 4 x i32> zeroinitializer
+  %b.insert = insertelement <vscale x 4 x float> poison, float %b, i64 0
+  %b.splat = shufflevector <vscale x 4 x float> %b.insert, <vscale x 4 x float> poison, <vscale x 4 x i32> zeroinitializer
+  %res = call <vscale x 8 x float> @llvm.vector.interleave2.nxv8f32(<vscale x 4 x float> %a.splat, <vscale x 4 x float> %b.splat)
+  ret <vscale x 8 x float> %res
+}
+
+declare <vscale x 2 x float> @llvm.vector.interleave2.nxv2f32(<vscale x 1 x float>, <vscale x 1 x float>)
+declare <vscale x 4 x float> @llvm.vector.interleave2.nxv4f32(<vscale x 2 x float>, <vscale x 2 x float>)
+declare <vscale x 8 x float> @llvm.vector.interleave2.nxv8f32(<vscale x 4 x float>, <vscale x 4 x float>)
+
+attributes #0 = { vscale_range(1,16) "target-features"="+sve" }



More information about the llvm-commits mailing list