[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