[llvm] [AArch64][GlobalISel] Add support for shuffle(v, undef) -> trn(v, v) transformation (PR #220914)
via llvm-commits
llvm-commits at lists.llvm.org
Thu Sep 3 07:08:07 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-backend-aarch64
Author: Joshua Rodriguez (JoshdRod)
<details>
<summary>Changes</summary>
Stacked PR: 3/3. Preceded by https://github.com/llvm/llvm-project/pull/220535.
In SDAG, the aarch64-isel phase checks if vector shuffles can be expressed as trns. To do this, it checks shuffles of type shuffle(v, v), and shuffle(v, undefined).
GlobalISel previously only checked shuffles of type shuffle(v, v). Add a check for the situation where one of the operands is undefined.
Notes:
A trn takes two vectors, places the even-indexed elements in the bottom half, and the odd-indexed elements in the top half.
e.g: `trn <0, 1, 2, 3>, <4, 5, 6, 7> => <0, 4, 2, 6, 1, 5, 3, 7>`.
A trn1 takes the bottom half of the result (aka. the even-indexed elements)
e.g: `trn1 <0, 1, 2, 3>, <4, 5, 6, 7> => <0, 4, 2, 6>`.
A shuffle is a generic MIR opcode that takes 2 vectors, and places elements of each into a single vector. The elements are selected based on a mask.
e.g: `G_SHUFFLE_VECTOR <A, B, C, D>, <E, F, G, H>, <0, 4, 6, 3> => <A, E, G, D>`.
For certain masks, a shuffle can be represented as a trn.
e.g: `G_SHUFFLE_VECTOR <A, B, C, D>, <A, B, C, D>, <0, 4, 2, 6> => trn1 <A, A, C, C>`.
---
Full diff: https://github.com/llvm/llvm-project/pull/220914.diff
4 Files Affected:
- (modified) llvm/lib/Target/AArch64/AArch64ISelLowering.cpp (+3-19)
- (modified) llvm/lib/Target/AArch64/AArch64PerfectShuffle.h (+16)
- (modified) llvm/lib/Target/AArch64/GISel/AArch64PostLegalizerLowering.cpp (+3-2)
- (modified) llvm/test/CodeGen/AArch64/arm64-trn.ll (+65-7)
``````````diff
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index 22198cd122fc7..6ba59504194ae 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -15518,22 +15518,6 @@ static bool isEXTMaskWithSplat(ArrayRef<int> M, EVT VT, unsigned SplatOperand,
return false;
}
-/// isTRN_v_undef_Mask - Special case of isTRNMask for canonical form of
-/// "vector_shuffle v, v", i.e., "vector_shuffle v, undef".
-/// Mask is e.g., <0, 0, 2, 2> instead of <0, 4, 2, 6>.
-static bool isTRN_v_undef_Mask(ArrayRef<int> M, EVT VT, unsigned &WhichResult) {
- unsigned NumElts = VT.getVectorNumElements();
- if (NumElts % 2 != 0)
- return false;
- WhichResult = (M[0] == 0 ? 0 : 1);
- for (unsigned i = 0; i < NumElts; i += 2) {
- if ((M[i] >= 0 && (unsigned)M[i] != i + WhichResult) ||
- (M[i + 1] >= 0 && (unsigned)M[i + 1] != i + WhichResult))
- return false;
- }
- return true;
-}
-
static bool isINSMask(ArrayRef<int> M, int NumInputElements,
bool &DstIsLeft, int &Anomaly) {
if (M.size() != static_cast<size_t>(NumInputElements))
@@ -16255,7 +16239,7 @@ SDValue AArch64TargetLowering::LowerVECTOR_SHUFFLE(SDValue Op,
unsigned Opc = (WhichResult == 0) ? AArch64ISD::UZP1 : AArch64ISD::UZP2;
return DAG.getNode(Opc, DL, V1.getValueType(), V1, V1);
}
- if (isTRN_v_undef_Mask(ShuffleMask, VT, WhichResult)) {
+ if (isTRN_v_undef_Mask(ShuffleMask, NumElts, WhichResult)) {
unsigned Opc = (WhichResult == 0) ? AArch64ISD::TRN1 : AArch64ISD::TRN2;
return DAG.getNode(Opc, DL, V1.getValueType(), V1, V1);
}
@@ -18038,7 +18022,7 @@ bool AArch64TargetLowering::isShuffleMaskLegal(ArrayRef<int> M, EVT VT) const {
isTRNMask(M, NumElts, DummyUnsigned, DummyUnsigned) ||
isUZPMask(M, NumElts, DummyUnsigned) ||
isZIPMask(M, NumElts, DummyUnsigned, DummyUnsigned) ||
- isTRN_v_undef_Mask(M, VT, DummyUnsigned) ||
+ isTRN_v_undef_Mask(M, NumElts, DummyUnsigned) ||
isUZP_v_undef_Mask(M, NumElts, DummyUnsigned) ||
isZIP_v_undef_Mask(M, NumElts, DummyUnsigned) ||
isINSMask(M, NumElts, DummyBool, DummyInt) ||
@@ -35618,7 +35602,7 @@ SDValue AArch64TargetLowering::LowerFixedLengthVECTOR_SHUFFLEToSVE(
return convertFromScalableVector(
DAG, VT, DAG.getNode(AArch64ISD::ZIP1, DL, ContainerVT, Op1, Op1));
- if (isTRN_v_undef_Mask(ShuffleMask, VT, WhichResult)) {
+ if (isTRN_v_undef_Mask(ShuffleMask, NumElts, WhichResult)) {
unsigned Opc = (WhichResult == 0) ? AArch64ISD::TRN1 : AArch64ISD::TRN2;
return convertFromScalableVector(
DAG, VT, DAG.getNode(Opc, DL, ContainerVT, Op1, Op1));
diff --git a/llvm/lib/Target/AArch64/AArch64PerfectShuffle.h b/llvm/lib/Target/AArch64/AArch64PerfectShuffle.h
index 9683248b2e66b..00d9edbbc02ea 100644
--- a/llvm/lib/Target/AArch64/AArch64PerfectShuffle.h
+++ b/llvm/lib/Target/AArch64/AArch64PerfectShuffle.h
@@ -221,6 +221,22 @@ inline bool isTRNMask(ArrayRef<int> M, unsigned NumElts,
return true;
}
+/// isTRN_v_undef_Mask - Special case of isTRNMask for canonical form of
+/// "vector_shuffle v, v", i.e., "vector_shuffle v, undef".
+/// Mask is e.g., <0, 0, 2, 2> instead of <0, 4, 2, 6>.
+inline bool isTRN_v_undef_Mask(ArrayRef<int> M, unsigned NumElts,
+ unsigned &WhichResult) {
+ if (NumElts % 2 != 0)
+ return false;
+ WhichResult = (M[0] == 0 ? 0 : 1);
+ for (unsigned i = 0; i < NumElts; i += 2) {
+ if ((M[i] >= 0 && (unsigned)M[i] != i + WhichResult) ||
+ (M[i + 1] >= 0 && (unsigned)M[i + 1] != i + WhichResult))
+ return false;
+ }
+ return true;
+}
+
/// isREVMask - Check if a vector shuffle corresponds to a REV
/// instruction with the specified blocksize. (The order of the elements
/// within each block of the vector is reversed.)
diff --git a/llvm/lib/Target/AArch64/GISel/AArch64PostLegalizerLowering.cpp b/llvm/lib/Target/AArch64/GISel/AArch64PostLegalizerLowering.cpp
index 283e0cfbebfe1..6256d78da324d 100644
--- a/llvm/lib/Target/AArch64/GISel/AArch64PostLegalizerLowering.cpp
+++ b/llvm/lib/Target/AArch64/GISel/AArch64PostLegalizerLowering.cpp
@@ -192,11 +192,12 @@ bool matchTRN(MachineInstr &MI, MachineRegisterInfo &MRI,
ShuffleVectorPseudo &MatchInfo) {
assert(MI.getOpcode() == TargetOpcode::G_SHUFFLE_VECTOR);
unsigned WhichResult;
- unsigned OperandOrder;
+ unsigned OperandOrder = 0;
ArrayRef<int> ShuffleMask = MI.getOperand(3).getShuffleMask();
Register Dst = MI.getOperand(0).getReg();
unsigned NumElts = MRI.getType(Dst).getNumElements();
- if (!isTRNMask(ShuffleMask, NumElts, WhichResult, OperandOrder))
+ if (!isTRNMask(ShuffleMask, NumElts, WhichResult, OperandOrder) &&
+ !isTRN_v_undef_Mask(ShuffleMask, NumElts, WhichResult))
return false;
unsigned Opc = (WhichResult == 0) ? AArch64::G_TRN1 : AArch64::G_TRN2;
Register V1 = MI.getOperand(OperandOrder == 0 ? 1 : 2).getReg();
diff --git a/llvm/test/CodeGen/AArch64/arm64-trn.ll b/llvm/test/CodeGen/AArch64/arm64-trn.ll
index 85a56c042136c..aba55be1b99ab 100644
--- a/llvm/test/CodeGen/AArch64/arm64-trn.ll
+++ b/llvm/test/CodeGen/AArch64/arm64-trn.ll
@@ -1,5 +1,6 @@
; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py
-; RUN: llc < %s -mtriple=aarch64 | FileCheck %s --check-prefixes=CHECKLE
+; RUN: llc < %s -mtriple=aarch64 | FileCheck %s --check-prefixes=CHECKLE,CHECKLE-SD
+; RUN: llc < %s -mtriple=aarch64 --global-isel --global-isel-abort=1 | FileCheck %s --check-prefixes=CHECKLE,CHECKLE-GI
; RUN: llc < %s -mtriple=aarch64_be | FileCheck %s --check-prefixes=CHECKBE
define <8 x i8> @vtrni8(ptr %A, ptr %B) nounwind {
@@ -57,12 +58,23 @@ define <4 x i16> @vtrni16(ptr %A, ptr %B) nounwind {
}
define <8 x i8> @vtrni16_viabitcast(ptr %A, ptr %B) nounwind {
-; CHECKLE-LABEL: vtrni16_viabitcast:
-; CHECKLE: // %bb.0:
-; CHECKLE-NEXT: ldr d0, [x0]
-; CHECKLE-NEXT: ldr d1, [x1]
-; CHECKLE-NEXT: trn1 v0.4h, v0.4h, v1.4h
-; CHECKLE-NEXT: ret
+; CHECKLE-SD-LABEL: vtrni16_viabitcast:
+; CHECKLE-SD: // %bb.0:
+; CHECKLE-SD-NEXT: ldr d0, [x0]
+; CHECKLE-SD-NEXT: ldr d1, [x1]
+; CHECKLE-SD-NEXT: trn1 v0.4h, v0.4h, v1.4h
+; CHECKLE-SD-NEXT: ret
+;
+; CHECKLE-GI-LABEL: vtrni16_viabitcast:
+; CHECKLE-GI: // %bb.0:
+; CHECKLE-GI-NEXT: ldr d0, [x0]
+; CHECKLE-GI-NEXT: ldr d1, [x1]
+; CHECKLE-GI-NEXT: adrp x8, .LCPI2_0
+; CHECKLE-GI-NEXT: mov v0.d[1], v1.d[0]
+; CHECKLE-GI-NEXT: ldr d1, [x8, :lo12:.LCPI2_0]
+; CHECKLE-GI-NEXT: tbl v0.16b, { v0.16b }, v1.16b
+; CHECKLE-GI-NEXT: // kill: def $d0 killed $d0 killed $q0
+; CHECKLE-GI-NEXT: ret
;
; CHECKBE-LABEL: vtrni16_viabitcast:
; CHECKBE: // %bb.0:
@@ -463,3 +475,49 @@ define <16 x i8> @vtrnQi8_undef_012(ptr %A, ptr %B) nounwind {
%tmp5 = add <16 x i8> %tmp3, %tmp4
ret <16 x i8> %tmp5
}
+
+define <8 x i8> @vtrnDi8_poison(<8 x i8> %a) {
+; CHECKLE-LABEL: vtrnDi8_poison:
+; CHECKLE: // %bb.0:
+; CHECKLE-NEXT: trn1 v1.8b, v0.8b, v0.8b
+; CHECKLE-NEXT: trn2 v0.8b, v0.8b, v0.8b
+; CHECKLE-NEXT: eor v0.8b, v1.8b, v0.8b
+; CHECKLE-NEXT: ret
+;
+; CHECKBE-LABEL: vtrnDi8_poison:
+; CHECKBE: // %bb.0:
+; CHECKBE-NEXT: rev64 v0.8b, v0.8b
+; CHECKBE-NEXT: trn1 v1.8b, v0.8b, v0.8b
+; CHECKBE-NEXT: trn2 v0.8b, v0.8b, v0.8b
+; CHECKBE-NEXT: eor v0.8b, v1.8b, v0.8b
+; CHECKBE-NEXT: rev64 v0.8b, v0.8b
+; CHECKBE-NEXT: ret
+ %tmp3 = shufflevector <8 x i8> %a, <8 x i8> poison, <8 x i32> <i32 0, i32 0, i32 2, i32 2, i32 4, i32 4, i32 6, i32 6>
+ %tmp4 = shufflevector <8 x i8> %a, <8 x i8> poison, <8 x i32> <i32 1, i32 1, i32 3, i32 3, i32 5, i32 5, i32 7, i32 7>
+ %ret = xor <8 x i8> %tmp3, %tmp4
+ ret <8 x i8> %ret
+}
+
+define <8 x i16> @vtrnQi8_poison(<8 x i16> %a) {
+; CHECKLE-LABEL: vtrnQi8_poison:
+; CHECKLE: // %bb.0:
+; CHECKLE-NEXT: trn1 v1.8h, v0.8h, v0.8h
+; CHECKLE-NEXT: trn2 v0.8h, v0.8h, v0.8h
+; CHECKLE-NEXT: eor v0.16b, v1.16b, v0.16b
+; CHECKLE-NEXT: ret
+;
+; CHECKBE-LABEL: vtrnQi8_poison:
+; CHECKBE: // %bb.0:
+; CHECKBE-NEXT: rev64 v0.8h, v0.8h
+; CHECKBE-NEXT: ext v0.16b, v0.16b, v0.16b, #8
+; CHECKBE-NEXT: trn1 v1.8h, v0.8h, v0.8h
+; CHECKBE-NEXT: trn2 v0.8h, v0.8h, v0.8h
+; CHECKBE-NEXT: eor v0.16b, v1.16b, v0.16b
+; CHECKBE-NEXT: rev64 v0.8h, v0.8h
+; CHECKBE-NEXT: ext v0.16b, v0.16b, v0.16b, #8
+; CHECKBE-NEXT: ret
+ %tmp3 = shufflevector <8 x i16> %a, <8 x i16> poison, <8 x i32> <i32 0, i32 0, i32 2, i32 2, i32 4, i32 4, i32 6, i32 6>
+ %tmp4 = shufflevector <8 x i16> %a, <8 x i16> poison, <8 x i32> <i32 1, i32 1, i32 3, i32 3, i32 5, i32 5, i32 7, i32 7>
+ %ret = xor <8 x i16> %tmp3, %tmp4
+ ret <8 x i16> %ret
+}
``````````
</details>
https://github.com/llvm/llvm-project/pull/220914
More information about the llvm-commits
mailing list