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

via llvm-commits llvm-commits at lists.llvm.org
Thu Aug 27 01:46:55 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-backend-aarch64

Author: Serval MARTINOT-LAGARDE (Serval6)

<details>
<summary>Changes</summary>

## Description
 
For narrow scalable vector types (e.g. `<vscale x 1 x T>`), interleaving two
splat vectors via `llvm.vector.interleave2` fails to legalize: the type
legalizer does not know how to widen the resulting `VECTOR_INTERLEAVE` node,
and `llc` crashes with:
 
```
LLVM ERROR: Do not know how to widen the result of this operator!
```
 
This is reproduced by the `interleave2_nxv2f32` / `interleave2_nxv4f32` cases
in the attached test, where two `<vscale x 1 x float>` splats are interleaved
into a `<vscale x 2/4 x float>` result.
 
This patch adds two AArch64-specific DAG combines on `concat_vectors`, which
recognize the interleave-of-splats pattern *before* the illegal intermediate
type reaches the widening legalizer:
 
- **`simplifyAlternatingMask`**: matches
 
```
  concat_vectors(vector_interleave(splat_vector(i1 0), splat_vector(i1 -1)))
```
 
  and folds it directly into a `PTRUE` (or its complement via `getNOT`) built
  at twice the element count, producing an alternating true/false predicate
  mask.
 
- **`simplifyAlternatingSplat`**: matches the general case
 
```
  concat_vectors(vector_interleave(splat_vector(A), splat_vector(B)))
```
 
  for legal, non-i1/i8/i16 vector types, and lowers it to
  `AArch64ISD::DUP_MERGE_PASSTHRU` driven by a half-element-count `PTRUE`.
 
Both combines are wired into `performConcatVectorsCombine`, in the
`VT.isScalableVector()` branch.
 
### Testing
 
Added `llvm/test/CodeGen/AArch64/sve-interleave-of-splat.ll` with three cases
(`nxv2f32`, `nxv4f32`, `nxv8f32`):
 
- The first two exercise the new combine and previously triggered the
  widening crash.
- The third is wide enough to be legal on its own and serves as a
  non-regression check, still lowering to `zip1`/`zip2`.
 
Verified that `llc` crashes on this test without the patch
(`DAGTypeLegalizer::WidenVectorResult` failing on the `vector_interleave` of
`<vscale x 1 x float>` splats in `@<!-- -->foo2`), and passes with the patch applied.


## Error before patch on the added test

```
# RUN: at line 2
llvm-project/build/bin/llc < llvm-project/llvm/test/CodeGen/AArch64/sve-interleave-of-splat.ll | llvm-project/build/bin/FileCheck llvm-project/llvm/test/CodeGen/AArch64/sve-interleave-of-splat.ll
# executed command: llvm-project/build/bin/llc
# .---command stderr------------
# | WidenVectorResult #<!-- -->0: t11: nxv1f32,nxv1f32 = vector_interleave t8, t10
# |
# | LLVM ERROR: Do not know how to widen the result of this operator!
# | PLEASE submit a bug report to https://github.com/llvm/llvm-project/issues/ and include the crash backtrace and instructions to reproduce the bug.
# | Stack dump:
# | 0.  Program arguments: llvm-project/build/bin/llc
# | 1.  Running pass 'Function Pass Manager' on module '<stdin>'.
# | 2.  Running pass 'AArch64 Instruction Selection' on function '@<!-- -->foo2'
# |  #<!-- -->0 0x00000000020dad30 llvm::sys::PrintStackTrace(llvm::raw_ostream&, int) (llvm-project/build/bin/llc+0x20dad30)
# |  #<!-- -->1 0x00000000020d822c llvm::sys::RunSignalHandlers() (llvm-project/build/bin/llc+0x20d822c)
# |  #<!-- -->2 0x00000000020d83a4 SignalHandler(int, siginfo_t*, void*) Signals.cpp:0:0
# |  #<!-- -->3 0x0000ffffbd6467f0 (linux-vdso.so.1+0x7f0)
# |  #<!-- -->4 0x0000ffffbd0d4bc8 __pthread_kill_implementation (/lib64/libc.so.6+0x82bc8)
# |  #<!-- -->5 0x0000ffffbd08ccbc gsignal (/lib64/libc.so.6+0x3acbc)
# |  #<!-- -->6 0x0000ffffbd079274 abort (/lib64/libc.so.6+0x27274)
# |  #<!-- -->7 0x000000000202d510 llvm::report_fatal_error(llvm::StringRef, bool) (llvm-project/build/bin/llc+0x202d510)
# |  #<!-- -->8 0x000000000202d570 llvm::reportFatalInternalError(char const*) (llvm-project/build/bin/llc+0x202d570)
# |  #<!-- -->9 0x0000000001f7aab8 llvm::DAGTypeLegalizer::WidenVectorResult(llvm::SDNode*, unsigned int) (llvm-project/build/bin/llc+0x1f7aab8)
# | #<!-- -->10 0x0000000001f28758 llvm::DAGTypeLegalizer::run() (llvm-project/build/bin/llc+0x1f28758)
# | #<!-- -->11 0x0000000001f28e90 llvm::SelectionDAG::LegalizeTypes() (llvm-project/build/bin/llc+0x1f28e90)
# | #<!-- -->12 0x0000000001e9c404 llvm::SelectionDAGISel::CodeGenAndEmitDAG() (llvm-project/build/bin/llc+0x1e9c404)
# | #<!-- -->13 0x0000000001e9f83c llvm::SelectionDAGISel::SelectAllBasicBlocks(llvm::Function const&) (llvm-project/build/bin/llc+0x1e9f83c)
# | #<!-- -->14 0x0000000001ea0ee4 llvm::SelectionDAGISel::runOnMachineFunction(llvm::MachineFunction&) (llvm-project/build/bin/llc+0x1ea0ee4)
# | #<!-- -->15 0x0000000001e9b328 llvm::SelectionDAGISelLegacy::runOnMachineFunction(llvm::MachineFunction&) (llvm-project/build/bin/llc+0x1e9b328)
# | #<!-- -->16 0x0000000001121bd0 llvm::MachineFunctionPass::runOnFunction(llvm::Function&) (.part.0) MachineFunctionPass.cpp:0:0
# | #<!-- -->17 0x00000000016d8b9c llvm::FPPassManager::runOnFunction(llvm::Function&) (llvm-project/build/bin/llc+0x16d8b9c)
# | #<!-- -->18 0x00000000016d8fc0 llvm::FPPassManager::runOnModule(llvm::Module&) (llvm-project/build/bin/llc+0x16d8fc0)
# | #<!-- -->19 0x00000000016d9998 (anonymous namespace)::MPPassManager::runOnModule(llvm::Module&) LegacyPassManager.cpp:0:0
# | #<!-- -->20 0x00000000016d9e64 llvm::legacy::PassManagerImpl::run(llvm::Module&) (llvm-project/build/bin/llc+0x16d9e64)
# | #<!-- -->21 0x00000000007b585c compileModule(char**, llvm::SmallVectorImpl<llvm::PassPlugin>&, llvm::LLVMContext&, std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char>>&) llc.cpp:0:0
# | #<!-- -->22 0x0000000000742ee4 main (llvm-project/build/bin/llc+0x742ee4)
# | #<!-- -->23 0x0000ffffbd079540 __libc_start_call_main (/lib64/libc.so.6+0x27540)
# | #<!-- -->24 0x0000ffffbd079618 __libc_start_main@<!-- -->GLIBC_2.17 (/lib64/libc.so.6+0x27618)
# | #<!-- -->25 0x00000000007ac170 _start (llvm-project/build/bin/llc+0x7ac170)
# `-----------------------------
# error: command failed with exit status: -6
# executed command: llvm-project/build/bin/FileCheck llvm-project/llvm/test/CodeGen/AArch64/sve-interleave-of-splat.ll
# .---command stderr------------
# | FileCheck error: '<stdin>' is empty.
# | FileCheck command line:  llvm-project/build/bin/FileCheck llvm-project/llvm/test/CodeGen/AArch64/sve-interleave-of-splat.ll
# `-----------------------------
# error: command failed with exit status: 2
 
--
 
****************************************
Failed Tests (1):
  LLVM :: CodeGen/AArch64/sve-interleave-of-splat.ll
```

---
Full diff: https://github.com/llvm/llvm-project/pull/219146.diff


2 Files Affected:

- (modified) llvm/lib/Target/AArch64/AArch64ISelLowering.cpp (+108-1) 
- (added) llvm/test/CodeGen/AArch64/sve-interleave-of-splat.ll (+64) 


``````````diff
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" }

``````````

</details>


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


More information about the llvm-commits mailing list