[Mlir-commits] [mlir] cfb9a0b - [SVE ACLE] Add svcount support to llvm.aarch64.sve.psel intrinsic. (#217360)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Tue Aug 25 03:56:13 PDT 2026
Author: Paul Walker
Date: 2026-08-25T11:56:05+01:00
New Revision: cfb9a0ba0dcff9e4470681072314dec896440faf
URL: https://github.com/llvm/llvm-project/commit/cfb9a0ba0dcff9e4470681072314dec896440faf
DIFF: https://github.com/llvm/llvm-project/commit/cfb9a0ba0dcff9e4470681072314dec896440faf.diff
LOG: [SVE ACLE] Add svcount support to llvm.aarch64.sve.psel intrinsic. (#217360)
PSEL has genuine support for predicate-as-counter, so it seems fair for
its equivalent intrinsic to have matching support rather than indirect
support using the svbool variant.
Added:
Modified:
clang/lib/CodeGen/TargetBuiltins/ARM.cpp
clang/test/CodeGen/AArch64/sve2p1-intrinsics/acle_sve2p1_psel.c
clang/test/CodeGen/AArch64/sve2p1-intrinsics/acle_sve2p1_psel_svcount.c
llvm/include/llvm/IR/IntrinsicsAArch64.td
llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
llvm/test/CodeGen/AArch64/sve2-intrinsics-psel.ll
mlir/include/mlir/Dialect/ArmSVE/IR/ArmSVE.td
mlir/test/Target/LLVMIR/arm-sve.mlir
Removed:
################################################################################
diff --git a/clang/lib/CodeGen/TargetBuiltins/ARM.cpp b/clang/lib/CodeGen/TargetBuiltins/ARM.cpp
index 662ec6d549986..5cddeef5e0f37 100644
--- a/clang/lib/CodeGen/TargetBuiltins/ARM.cpp
+++ b/clang/lib/CodeGen/TargetBuiltins/ARM.cpp
@@ -4086,24 +4086,11 @@ Value *CodeGenFunction::EmitAArch64SVEBuiltinExpr(unsigned BuiltinID,
case SVE::BI__builtin_sve_svpsel_lane_c16:
case SVE::BI__builtin_sve_svpsel_lane_c32:
case SVE::BI__builtin_sve_svpsel_lane_c64: {
- bool IsSVCount = isa<TargetExtType>(Ops[0]->getType());
- assert(((!IsSVCount || cast<TargetExtType>(Ops[0]->getType())->getName() ==
- "aarch64.svcount")) &&
- "Unexpected TargetExtType");
- auto SVCountTy =
- llvm::TargetExtType::get(getLLVMContext(), "aarch64.svcount");
- Function *CastFromSVCountF =
- CGM.getIntrinsic(Intrinsic::aarch64_sve_convert_to_svbool, SVCountTy);
- Function *CastToSVCountF =
- CGM.getIntrinsic(Intrinsic::aarch64_sve_convert_from_svbool, SVCountTy);
-
auto OverloadedTy = getSVEType(SVETypeFlags(Builtin->TypeModifier));
- Function *F = CGM.getIntrinsic(Intrinsic::aarch64_sve_psel, OverloadedTy);
- llvm::Value *Ops0 =
- IsSVCount ? Builder.CreateCall(CastFromSVCountF, Ops[0]) : Ops[0];
+ Function *F = CGM.getIntrinsic(Intrinsic::aarch64_sve_psel,
+ {Ops[0]->getType(), OverloadedTy});
llvm::Value *Ops1 = EmitSVEPredicateCast(Ops[1], OverloadedTy);
- llvm::Value *PSel = Builder.CreateCall(F, {Ops0, Ops1, Ops[2]});
- return IsSVCount ? Builder.CreateCall(CastToSVCountF, PSel) : PSel;
+ return Builder.CreateCall(F, {Ops[0], Ops1, Ops[2]});
}
case SVE::BI__builtin_sve_svmov_b_z: {
// svmov_b_z(pg, op) <=> svand_b_z(pg, op, op)
diff --git a/clang/test/CodeGen/AArch64/sve2p1-intrinsics/acle_sve2p1_psel.c b/clang/test/CodeGen/AArch64/sve2p1-intrinsics/acle_sve2p1_psel.c
index 8b7761e981c96..1ee8daa2df4ad 100644
--- a/clang/test/CodeGen/AArch64/sve2p1-intrinsics/acle_sve2p1_psel.c
+++ b/clang/test/CodeGen/AArch64/sve2p1-intrinsics/acle_sve2p1_psel.c
@@ -26,13 +26,13 @@
// CHECK-LABEL: @test_svpsel_lane_b8(
// CHECK-NEXT: entry:
// CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 15
-// CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 16 x i1> [[P2:%.*]], i32 [[ADD]])
+// CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1.nxv16i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 16 x i1> [[P2:%.*]], i32 [[ADD]])
// CHECK-NEXT: ret <vscale x 16 x i1> [[TMP0]]
//
// CPP-CHECK-LABEL: @_Z19test_svpsel_lane_b8u10__SVBool_tS_j(
// CPP-CHECK-NEXT: entry:
// CPP-CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 15
-// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 16 x i1> [[P2:%.*]], i32 [[ADD]])
+// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1.nxv16i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 16 x i1> [[P2:%.*]], i32 [[ADD]])
// CPP-CHECK-NEXT: ret <vscale x 16 x i1> [[TMP0]]
//
svbool_t test_svpsel_lane_b8(svbool_t p1, svbool_t p2, uint32_t idx) ATTR {
@@ -43,14 +43,14 @@ svbool_t test_svpsel_lane_b8(svbool_t p1, svbool_t p2, uint32_t idx) ATTR {
// CHECK-NEXT: entry:
// CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 7
// CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 8 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv8i1(<vscale x 16 x i1> [[P2:%.*]])
-// CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv8i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 8 x i1> [[TMP0]], i32 [[ADD]])
+// CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1.nxv8i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 8 x i1> [[TMP0]], i32 [[ADD]])
// CHECK-NEXT: ret <vscale x 16 x i1> [[TMP1]]
//
// CPP-CHECK-LABEL: @_Z20test_svpsel_lane_b16u10__SVBool_tS_j(
// CPP-CHECK-NEXT: entry:
// CPP-CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 7
// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 8 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv8i1(<vscale x 16 x i1> [[P2:%.*]])
-// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv8i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 8 x i1> [[TMP0]], i32 [[ADD]])
+// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1.nxv8i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 8 x i1> [[TMP0]], i32 [[ADD]])
// CPP-CHECK-NEXT: ret <vscale x 16 x i1> [[TMP1]]
//
svbool_t test_svpsel_lane_b16(svbool_t p1, svbool_t p2, uint32_t idx) ATTR {
@@ -61,14 +61,14 @@ svbool_t test_svpsel_lane_b16(svbool_t p1, svbool_t p2, uint32_t idx) ATTR {
// CHECK-NEXT: entry:
// CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 3
// CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 4 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv4i1(<vscale x 16 x i1> [[P2:%.*]])
-// CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv4i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 4 x i1> [[TMP0]], i32 [[ADD]])
+// CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1.nxv4i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 4 x i1> [[TMP0]], i32 [[ADD]])
// CHECK-NEXT: ret <vscale x 16 x i1> [[TMP1]]
//
// CPP-CHECK-LABEL: @_Z20test_svpsel_lane_b32u10__SVBool_tS_j(
// CPP-CHECK-NEXT: entry:
// CPP-CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 3
// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 4 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv4i1(<vscale x 16 x i1> [[P2:%.*]])
-// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv4i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 4 x i1> [[TMP0]], i32 [[ADD]])
+// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1.nxv4i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 4 x i1> [[TMP0]], i32 [[ADD]])
// CPP-CHECK-NEXT: ret <vscale x 16 x i1> [[TMP1]]
//
svbool_t test_svpsel_lane_b32(svbool_t p1, svbool_t p2, uint32_t idx) ATTR {
@@ -79,14 +79,14 @@ svbool_t test_svpsel_lane_b32(svbool_t p1, svbool_t p2, uint32_t idx) ATTR {
// CHECK-NEXT: entry:
// CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 1
// CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 2 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv2i1(<vscale x 16 x i1> [[P2:%.*]])
-// CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv2i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 2 x i1> [[TMP0]], i32 [[ADD]])
+// CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1.nxv2i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 2 x i1> [[TMP0]], i32 [[ADD]])
// CHECK-NEXT: ret <vscale x 16 x i1> [[TMP1]]
//
// CPP-CHECK-LABEL: @_Z20test_svpsel_lane_b64u10__SVBool_tS_j(
// CPP-CHECK-NEXT: entry:
// CPP-CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 1
// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 2 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv2i1(<vscale x 16 x i1> [[P2:%.*]])
-// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv2i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 2 x i1> [[TMP0]], i32 [[ADD]])
+// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1.nxv2i1(<vscale x 16 x i1> [[P1:%.*]], <vscale x 2 x i1> [[TMP0]], i32 [[ADD]])
// CPP-CHECK-NEXT: ret <vscale x 16 x i1> [[TMP1]]
//
svbool_t test_svpsel_lane_b64(svbool_t p1, svbool_t p2, uint32_t idx) ATTR {
diff --git a/clang/test/CodeGen/AArch64/sve2p1-intrinsics/acle_sve2p1_psel_svcount.c b/clang/test/CodeGen/AArch64/sve2p1-intrinsics/acle_sve2p1_psel_svcount.c
index 589c540dcafb6..43383f1fae68c 100644
--- a/clang/test/CodeGen/AArch64/sve2p1-intrinsics/acle_sve2p1_psel_svcount.c
+++ b/clang/test/CodeGen/AArch64/sve2p1-intrinsics/acle_sve2p1_psel_svcount.c
@@ -26,18 +26,14 @@
// CHECK-LABEL: @test_svpsel_lane_c8(
// CHECK-NEXT: entry:
// CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 15
-// CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.convert.to.svbool.taarch64.svcountt(target("aarch64.svcount") [[P1:%.*]])
-// CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1(<vscale x 16 x i1> [[TMP0]], <vscale x 16 x i1> [[P2:%.*]], i32 [[ADD]])
-// CHECK-NEXT: [[TMP2:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.convert.from.svbool.taarch64.svcountt(<vscale x 16 x i1> [[TMP1]])
-// CHECK-NEXT: ret target("aarch64.svcount") [[TMP2]]
+// CHECK-NEXT: [[TMP0:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.psel.taarch64.svcountt.nxv16i1(target("aarch64.svcount") [[P1:%.*]], <vscale x 16 x i1> [[P2:%.*]], i32 [[ADD]])
+// CHECK-NEXT: ret target("aarch64.svcount") [[TMP0]]
//
// CPP-CHECK-LABEL: @_Z19test_svpsel_lane_c8u11__SVCount_tu10__SVBool_tj(
// CPP-CHECK-NEXT: entry:
// CPP-CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 15
-// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.convert.to.svbool.taarch64.svcountt(target("aarch64.svcount") [[P1:%.*]])
-// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1(<vscale x 16 x i1> [[TMP0]], <vscale x 16 x i1> [[P2:%.*]], i32 [[ADD]])
-// CPP-CHECK-NEXT: [[TMP2:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.convert.from.svbool.taarch64.svcountt(<vscale x 16 x i1> [[TMP1]])
-// CPP-CHECK-NEXT: ret target("aarch64.svcount") [[TMP2]]
+// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.psel.taarch64.svcountt.nxv16i1(target("aarch64.svcount") [[P1:%.*]], <vscale x 16 x i1> [[P2:%.*]], i32 [[ADD]])
+// CPP-CHECK-NEXT: ret target("aarch64.svcount") [[TMP0]]
//
svcount_t test_svpsel_lane_c8(svcount_t p1, svbool_t p2, uint32_t idx) ATTR {
return svpsel_lane_c8(p1, p2, idx + 15);
@@ -46,20 +42,16 @@ svcount_t test_svpsel_lane_c8(svcount_t p1, svbool_t p2, uint32_t idx) ATTR {
// CHECK-LABEL: @test_svpsel_lane_c16(
// CHECK-NEXT: entry:
// CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 7
-// CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.convert.to.svbool.taarch64.svcountt(target("aarch64.svcount") [[P1:%.*]])
-// CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 8 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv8i1(<vscale x 16 x i1> [[P2:%.*]])
-// CHECK-NEXT: [[TMP2:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv8i1(<vscale x 16 x i1> [[TMP0]], <vscale x 8 x i1> [[TMP1]], i32 [[ADD]])
-// CHECK-NEXT: [[TMP3:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.convert.from.svbool.taarch64.svcountt(<vscale x 16 x i1> [[TMP2]])
-// CHECK-NEXT: ret target("aarch64.svcount") [[TMP3]]
+// CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 8 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv8i1(<vscale x 16 x i1> [[P2:%.*]])
+// CHECK-NEXT: [[TMP1:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.psel.taarch64.svcountt.nxv8i1(target("aarch64.svcount") [[P1:%.*]], <vscale x 8 x i1> [[TMP0]], i32 [[ADD]])
+// CHECK-NEXT: ret target("aarch64.svcount") [[TMP1]]
//
// CPP-CHECK-LABEL: @_Z20test_svpsel_lane_c16u11__SVCount_tu10__SVBool_tj(
// CPP-CHECK-NEXT: entry:
// CPP-CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 7
-// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.convert.to.svbool.taarch64.svcountt(target("aarch64.svcount") [[P1:%.*]])
-// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 8 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv8i1(<vscale x 16 x i1> [[P2:%.*]])
-// CPP-CHECK-NEXT: [[TMP2:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv8i1(<vscale x 16 x i1> [[TMP0]], <vscale x 8 x i1> [[TMP1]], i32 [[ADD]])
-// CPP-CHECK-NEXT: [[TMP3:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.convert.from.svbool.taarch64.svcountt(<vscale x 16 x i1> [[TMP2]])
-// CPP-CHECK-NEXT: ret target("aarch64.svcount") [[TMP3]]
+// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 8 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv8i1(<vscale x 16 x i1> [[P2:%.*]])
+// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.psel.taarch64.svcountt.nxv8i1(target("aarch64.svcount") [[P1:%.*]], <vscale x 8 x i1> [[TMP0]], i32 [[ADD]])
+// CPP-CHECK-NEXT: ret target("aarch64.svcount") [[TMP1]]
//
svcount_t test_svpsel_lane_c16(svcount_t p1, svbool_t p2, uint32_t idx) ATTR {
return svpsel_lane_c16(p1, p2, idx + 7);
@@ -68,20 +60,16 @@ svcount_t test_svpsel_lane_c16(svcount_t p1, svbool_t p2, uint32_t idx) ATTR {
// CHECK-LABEL: @test_svpsel_lane_c32(
// CHECK-NEXT: entry:
// CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 3
-// CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.convert.to.svbool.taarch64.svcountt(target("aarch64.svcount") [[P1:%.*]])
-// CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 4 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv4i1(<vscale x 16 x i1> [[P2:%.*]])
-// CHECK-NEXT: [[TMP2:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv4i1(<vscale x 16 x i1> [[TMP0]], <vscale x 4 x i1> [[TMP1]], i32 [[ADD]])
-// CHECK-NEXT: [[TMP3:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.convert.from.svbool.taarch64.svcountt(<vscale x 16 x i1> [[TMP2]])
-// CHECK-NEXT: ret target("aarch64.svcount") [[TMP3]]
+// CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 4 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv4i1(<vscale x 16 x i1> [[P2:%.*]])
+// CHECK-NEXT: [[TMP1:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.psel.taarch64.svcountt.nxv4i1(target("aarch64.svcount") [[P1:%.*]], <vscale x 4 x i1> [[TMP0]], i32 [[ADD]])
+// CHECK-NEXT: ret target("aarch64.svcount") [[TMP1]]
//
// CPP-CHECK-LABEL: @_Z20test_svpsel_lane_c32u11__SVCount_tu10__SVBool_tj(
// CPP-CHECK-NEXT: entry:
// CPP-CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 3
-// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.convert.to.svbool.taarch64.svcountt(target("aarch64.svcount") [[P1:%.*]])
-// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 4 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv4i1(<vscale x 16 x i1> [[P2:%.*]])
-// CPP-CHECK-NEXT: [[TMP2:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv4i1(<vscale x 16 x i1> [[TMP0]], <vscale x 4 x i1> [[TMP1]], i32 [[ADD]])
-// CPP-CHECK-NEXT: [[TMP3:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.convert.from.svbool.taarch64.svcountt(<vscale x 16 x i1> [[TMP2]])
-// CPP-CHECK-NEXT: ret target("aarch64.svcount") [[TMP3]]
+// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 4 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv4i1(<vscale x 16 x i1> [[P2:%.*]])
+// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.psel.taarch64.svcountt.nxv4i1(target("aarch64.svcount") [[P1:%.*]], <vscale x 4 x i1> [[TMP0]], i32 [[ADD]])
+// CPP-CHECK-NEXT: ret target("aarch64.svcount") [[TMP1]]
//
svcount_t test_svpsel_lane_c32(svcount_t p1, svbool_t p2, uint32_t idx) ATTR {
return svpsel_lane_c32(p1, p2, idx + 3);
@@ -90,20 +78,16 @@ svcount_t test_svpsel_lane_c32(svcount_t p1, svbool_t p2, uint32_t idx) ATTR {
// CHECK-LABEL: @test_svpsel_lane_c64(
// CHECK-NEXT: entry:
// CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 1
-// CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.convert.to.svbool.taarch64.svcountt(target("aarch64.svcount") [[P1:%.*]])
-// CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 2 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv2i1(<vscale x 16 x i1> [[P2:%.*]])
-// CHECK-NEXT: [[TMP2:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv2i1(<vscale x 16 x i1> [[TMP0]], <vscale x 2 x i1> [[TMP1]], i32 [[ADD]])
-// CHECK-NEXT: [[TMP3:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.convert.from.svbool.taarch64.svcountt(<vscale x 16 x i1> [[TMP2]])
-// CHECK-NEXT: ret target("aarch64.svcount") [[TMP3]]
+// CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 2 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv2i1(<vscale x 16 x i1> [[P2:%.*]])
+// CHECK-NEXT: [[TMP1:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.psel.taarch64.svcountt.nxv2i1(target("aarch64.svcount") [[P1:%.*]], <vscale x 2 x i1> [[TMP0]], i32 [[ADD]])
+// CHECK-NEXT: ret target("aarch64.svcount") [[TMP1]]
//
// CPP-CHECK-LABEL: @_Z20test_svpsel_lane_c64u11__SVCount_tu10__SVBool_tj(
// CPP-CHECK-NEXT: entry:
// CPP-CHECK-NEXT: [[ADD:%.*]] = add i32 [[IDX:%.*]], 1
-// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.convert.to.svbool.taarch64.svcountt(target("aarch64.svcount") [[P1:%.*]])
-// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call <vscale x 2 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv2i1(<vscale x 16 x i1> [[P2:%.*]])
-// CPP-CHECK-NEXT: [[TMP2:%.*]] = tail call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv2i1(<vscale x 16 x i1> [[TMP0]], <vscale x 2 x i1> [[TMP1]], i32 [[ADD]])
-// CPP-CHECK-NEXT: [[TMP3:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.convert.from.svbool.taarch64.svcountt(<vscale x 16 x i1> [[TMP2]])
-// CPP-CHECK-NEXT: ret target("aarch64.svcount") [[TMP3]]
+// CPP-CHECK-NEXT: [[TMP0:%.*]] = tail call <vscale x 2 x i1> @llvm.aarch64.sve.convert.from.svbool.nxv2i1(<vscale x 16 x i1> [[P2:%.*]])
+// CPP-CHECK-NEXT: [[TMP1:%.*]] = tail call target("aarch64.svcount") @llvm.aarch64.sve.psel.taarch64.svcountt.nxv2i1(target("aarch64.svcount") [[P1:%.*]], <vscale x 2 x i1> [[TMP0]], i32 [[ADD]])
+// CPP-CHECK-NEXT: ret target("aarch64.svcount") [[TMP1]]
//
svcount_t test_svpsel_lane_c64(svcount_t p1, svbool_t p2, uint32_t idx) ATTR {
return svpsel_lane_c64(p1, p2, idx + 1);
diff --git a/llvm/include/llvm/IR/IntrinsicsAArch64.td b/llvm/include/llvm/IR/IntrinsicsAArch64.td
index 074f26cad2132..8fbe68f98cc63 100644
--- a/llvm/include/llvm/IR/IntrinsicsAArch64.td
+++ b/llvm/include/llvm/IR/IntrinsicsAArch64.td
@@ -3281,8 +3281,8 @@ let TargetPrefix = "aarch64" in {
//
def int_aarch64_sve_psel
- : DefaultAttrsIntrinsic<[llvm_nxv16i1_ty],
- [llvm_nxv16i1_ty,
+ : DefaultAttrsIntrinsic<[llvm_any_ty],
+ [LLVMMatchType<0>,
llvm_anyvector_ty, llvm_i32_ty],
[IntrNoMem]>;
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index 4ecdc6ba54b77..fc743256bc1a3 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -7301,6 +7301,15 @@ SDValue AArch64TargetLowering::LowerINTRINSIC_WO_CHAIN(SDValue Op,
case Intrinsic::aarch64_neon_pmul:
return DAG.getNode(ISD::CLMUL, DL, Op.getValueType(), Op.getOperand(1),
Op.getOperand(2));
+ case Intrinsic::aarch64_sve_psel: {
+ if (Op.getValueType() != MVT::aarch64svcount)
+ return Op;
+ // Lower predicate-as-counter variants to reuse normal predicate isel.
+ SmallVector<SDValue> Ops(Op->ops());
+ Ops[1] = DAG.getNode(ISD::BITCAST, DL, MVT::nxv16i1, Ops[1]);
+ SDValue PSel = DAG.getNode(ISD::INTRINSIC_WO_CHAIN, DL, MVT::nxv16i1, Ops);
+ return DAG.getNode(ISD::BITCAST, DL, MVT::aarch64svcount, PSel);
+ }
}
}
diff --git a/llvm/test/CodeGen/AArch64/sve2-intrinsics-psel.ll b/llvm/test/CodeGen/AArch64/sve2-intrinsics-psel.ll
index 01f7142831914..aaa7c4d582958 100644
--- a/llvm/test/CodeGen/AArch64/sve2-intrinsics-psel.ll
+++ b/llvm/test/CodeGen/AArch64/sve2-intrinsics-psel.ll
@@ -9,7 +9,7 @@ define <vscale x 16 x i1> @psel_b(<vscale x 16 x i1> %p1, <vscale x 16 x i1> %p2
; CHECK-NEXT: mov w12, w0
; CHECK-NEXT: psel p0, p0, p1.b[w12, 0]
; CHECK-NEXT: ret
- %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1(<vscale x 16 x i1> %p1, <vscale x 16 x i1> %p2, i32 %idx)
+ %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel(<vscale x 16 x i1> %p1, <vscale x 16 x i1> %p2, i32 %idx)
ret <vscale x 16 x i1> %res
}
@@ -20,7 +20,7 @@ define <vscale x 16 x i1> @psel_b_imm(<vscale x 16 x i1> %p1, <vscale x 16 x i1>
; CHECK-NEXT: psel p0, p0, p1.b[w12, 15]
; CHECK-NEXT: ret
%add = add i32 %idx, 15
- %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1(<vscale x 16 x i1> %p1, <vscale x 16 x i1> %p2, i32 %add)
+ %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel(<vscale x 16 x i1> %p1, <vscale x 16 x i1> %p2, i32 %add)
ret <vscale x 16 x i1> %res
}
@@ -30,7 +30,7 @@ define <vscale x 16 x i1> @psel_h(<vscale x 16 x i1> %p1, <vscale x 8 x i1> %p2,
; CHECK-NEXT: mov w12, w0
; CHECK-NEXT: psel p0, p0, p1.h[w12, 0]
; CHECK-NEXT: ret
- %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv8i1(<vscale x 16 x i1> %p1, <vscale x 8 x i1> %p2, i32 %idx)
+ %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel(<vscale x 16 x i1> %p1, <vscale x 8 x i1> %p2, i32 %idx)
ret <vscale x 16 x i1> %res
}
@@ -41,7 +41,7 @@ define <vscale x 16 x i1> @psel_h_imm(<vscale x 16 x i1> %p1, <vscale x 8 x i1>
; CHECK-NEXT: psel p0, p0, p1.h[w12, 7]
; CHECK-NEXT: ret
%add = add i32 %idx, 7
- %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv8i1(<vscale x 16 x i1> %p1, <vscale x 8 x i1> %p2, i32 %add)
+ %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel(<vscale x 16 x i1> %p1, <vscale x 8 x i1> %p2, i32 %add)
ret <vscale x 16 x i1> %res
}
@@ -51,7 +51,7 @@ define <vscale x 16 x i1> @psel_s(<vscale x 16 x i1> %p1, <vscale x 4 x i1> %p2,
; CHECK-NEXT: mov w12, w0
; CHECK-NEXT: psel p0, p0, p1.s[w12, 0]
; CHECK-NEXT: ret
- %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv4i1(<vscale x 16 x i1> %p1, <vscale x 4 x i1> %p2, i32 %idx)
+ %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel(<vscale x 16 x i1> %p1, <vscale x 4 x i1> %p2, i32 %idx)
ret <vscale x 16 x i1> %res
}
@@ -62,7 +62,7 @@ define <vscale x 16 x i1> @psel_s_imm(<vscale x 16 x i1> %p1, <vscale x 4 x i1>
; CHECK-NEXT: psel p0, p0, p1.s[w12, 3]
; CHECK-NEXT: ret
%add = add i32 %idx, 3
- %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv4i1(<vscale x 16 x i1> %p1, <vscale x 4 x i1> %p2, i32 %add)
+ %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel(<vscale x 16 x i1> %p1, <vscale x 4 x i1> %p2, i32 %add)
ret <vscale x 16 x i1> %res
}
@@ -72,7 +72,7 @@ define <vscale x 16 x i1> @psel_d(<vscale x 16 x i1> %p1, <vscale x 2 x i1> %p2,
; CHECK-NEXT: mov w12, w0
; CHECK-NEXT: psel p0, p0, p1.d[w12, 0]
; CHECK-NEXT: ret
- %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv2i1(<vscale x 16 x i1> %p1, <vscale x 2 x i1> %p2, i32 %idx)
+ %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel(<vscale x 16 x i1> %p1, <vscale x 2 x i1> %p2, i32 %idx)
ret <vscale x 16 x i1> %res
}
@@ -83,11 +83,90 @@ define <vscale x 16 x i1> @psel_d_imm(<vscale x 16 x i1> %p1, <vscale x 2 x i1>
; CHECK-NEXT: psel p0, p0, p1.d[w12, 1]
; CHECK-NEXT: ret
%add = add i32 %idx, 1
- %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv2i1(<vscale x 16 x i1> %p1, <vscale x 2 x i1> %p2, i32 %add)
+ %res = call <vscale x 16 x i1> @llvm.aarch64.sve.psel(<vscale x 16 x i1> %p1, <vscale x 2 x i1> %p2, i32 %add)
ret <vscale x 16 x i1> %res
}
-declare <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1(<vscale x 16 x i1>, <vscale x 16 x i1>, i32)
-declare <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv8i1(<vscale x 16 x i1>, <vscale x 8 x i1>, i32)
-declare <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv4i1(<vscale x 16 x i1>, <vscale x 4 x i1>, i32)
-declare <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv2i1(<vscale x 16 x i1>, <vscale x 2 x i1>, i32)
+define target("aarch64.svcount") @psel_svcount_b(target("aarch64.svcount") %p1, <vscale x 16 x i1> %p2, i32 %idx) {
+; CHECK-LABEL: psel_svcount_b:
+; CHECK: // %bb.0:
+; CHECK-NEXT: mov w12, w0
+; CHECK-NEXT: psel p0, p0, p1.b[w12, 0]
+; CHECK-NEXT: ret
+ %res = call target("aarch64.svcount") @llvm.aarch64.sve.psel(target("aarch64.svcount") %p1, <vscale x 16 x i1> %p2, i32 %idx)
+ ret target("aarch64.svcount") %res
+}
+
+define target("aarch64.svcount") @psel_svcount_b_imm(target("aarch64.svcount") %p1, <vscale x 16 x i1> %p2, i32 %idx) {
+; CHECK-LABEL: psel_svcount_b_imm:
+; CHECK: // %bb.0:
+; CHECK-NEXT: mov w12, w0
+; CHECK-NEXT: psel p0, p0, p1.b[w12, 15]
+; CHECK-NEXT: ret
+ %add = add i32 %idx, 15
+ %res = call target("aarch64.svcount") @llvm.aarch64.sve.psel(target("aarch64.svcount") %p1, <vscale x 16 x i1> %p2, i32 %add)
+ ret target("aarch64.svcount") %res
+}
+
+define target("aarch64.svcount") @psel_svcount_h(target("aarch64.svcount") %p1, <vscale x 8 x i1> %p2, i32 %idx) {
+; CHECK-LABEL: psel_svcount_h:
+; CHECK: // %bb.0:
+; CHECK-NEXT: mov w12, w0
+; CHECK-NEXT: psel p0, p0, p1.h[w12, 0]
+; CHECK-NEXT: ret
+ %res = call target("aarch64.svcount") @llvm.aarch64.sve.psel(target("aarch64.svcount") %p1, <vscale x 8 x i1> %p2, i32 %idx)
+ ret target("aarch64.svcount") %res
+}
+
+define target("aarch64.svcount") @psel_svcount_h_imm(target("aarch64.svcount") %p1, <vscale x 8 x i1> %p2, i32 %idx) {
+; CHECK-LABEL: psel_svcount_h_imm:
+; CHECK: // %bb.0:
+; CHECK-NEXT: mov w12, w0
+; CHECK-NEXT: psel p0, p0, p1.h[w12, 7]
+; CHECK-NEXT: ret
+ %add = add i32 %idx, 7
+ %res = call target("aarch64.svcount") @llvm.aarch64.sve.psel(target("aarch64.svcount") %p1, <vscale x 8 x i1> %p2, i32 %add)
+ ret target("aarch64.svcount") %res
+}
+
+define target("aarch64.svcount") @psel_svcount_s(target("aarch64.svcount") %p1, <vscale x 4 x i1> %p2, i32 %idx) {
+; CHECK-LABEL: psel_svcount_s:
+; CHECK: // %bb.0:
+; CHECK-NEXT: mov w12, w0
+; CHECK-NEXT: psel p0, p0, p1.s[w12, 0]
+; CHECK-NEXT: ret
+ %res = call target("aarch64.svcount") @llvm.aarch64.sve.psel(target("aarch64.svcount") %p1, <vscale x 4 x i1> %p2, i32 %idx)
+ ret target("aarch64.svcount") %res
+}
+
+define target("aarch64.svcount") @psel_svcount_s_imm(target("aarch64.svcount") %p1, <vscale x 4 x i1> %p2, i32 %idx) {
+; CHECK-LABEL: psel_svcount_s_imm:
+; CHECK: // %bb.0:
+; CHECK-NEXT: mov w12, w0
+; CHECK-NEXT: psel p0, p0, p1.s[w12, 3]
+; CHECK-NEXT: ret
+ %add = add i32 %idx, 3
+ %res = call target("aarch64.svcount") @llvm.aarch64.sve.psel(target("aarch64.svcount") %p1, <vscale x 4 x i1> %p2, i32 %add)
+ ret target("aarch64.svcount") %res
+}
+
+define target("aarch64.svcount") @psel_svcount_d(target("aarch64.svcount") %p1, <vscale x 2 x i1> %p2, i32 %idx) {
+; CHECK-LABEL: psel_svcount_d:
+; CHECK: // %bb.0:
+; CHECK-NEXT: mov w12, w0
+; CHECK-NEXT: psel p0, p0, p1.d[w12, 0]
+; CHECK-NEXT: ret
+ %res = call target("aarch64.svcount") @llvm.aarch64.sve.psel(target("aarch64.svcount") %p1, <vscale x 2 x i1> %p2, i32 %idx)
+ ret target("aarch64.svcount") %res
+}
+
+define target("aarch64.svcount") @psel_svcount_d_imm(target("aarch64.svcount") %p1, <vscale x 2 x i1> %p2, i32 %idx) {
+; CHECK-LABEL: psel_svcount_d_imm:
+; CHECK: // %bb.0:
+; CHECK-NEXT: mov w12, w0
+; CHECK-NEXT: psel p0, p0, p1.d[w12, 1]
+; CHECK-NEXT: ret
+ %add = add i32 %idx, 1
+ %res = call target("aarch64.svcount") @llvm.aarch64.sve.psel(target("aarch64.svcount") %p1, <vscale x 2 x i1> %p2, i32 %add)
+ ret target("aarch64.svcount") %res
+}
diff --git a/mlir/include/mlir/Dialect/ArmSVE/IR/ArmSVE.td b/mlir/include/mlir/Dialect/ArmSVE/IR/ArmSVE.td
index 61bb13b76167c..be4d9b123ac82 100644
--- a/mlir/include/mlir/Dialect/ArmSVE/IR/ArmSVE.td
+++ b/mlir/include/mlir/Dialect/ArmSVE/IR/ArmSVE.td
@@ -709,7 +709,7 @@ def ZipX4IntrOp : ArmSVE_IntrOp<"zip.x4",
// Note: This intrinsic requires SME or SVE2.1.
def PselIntrOp : ArmSVE_IntrOp<"psel",
/*traits=*/[Pure, TypeIs<"res", SVBool>],
- /*overloadedOperands=*/[1]>,
+ /*overloadedOperands=*/[0,1]>,
Arguments<(ins Arg<SVBool, "p1">:$p1,
Arg<SVEPredicate, "p2">:$p2,
Arg<I32, "index">:$index)>;
diff --git a/mlir/test/Target/LLVMIR/arm-sve.mlir b/mlir/test/Target/LLVMIR/arm-sve.mlir
index e0c4c48c5e9ab..67f383ea1a8a0 100644
--- a/mlir/test/Target/LLVMIR/arm-sve.mlir
+++ b/mlir/test/Target/LLVMIR/arm-sve.mlir
@@ -404,13 +404,13 @@ llvm.func @arm_sve_whilelt(%base: i64, %n: i64) {
// CHECK-SAME: <vscale x 16 x i1> %[[P4:[0-9]+]],
// CHECK-SAME: i32 %[[INDEX:[0-9]+]])
llvm.func @arm_sve_psel(%pn: vector<[16]xi1>, %p1: vector<[2]xi1>, %p2: vector<[4]xi1>, %p3: vector<[8]xi1>, %p4: vector<[16]xi1>, %index: i32) {
- // CHECK: call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv2i1(<vscale x 16 x i1> %[[PN]], <vscale x 2 x i1> %[[P1]], i32 %[[INDEX]])
+ // CHECK: call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1.nxv2i1(<vscale x 16 x i1> %[[PN]], <vscale x 2 x i1> %[[P1]], i32 %[[INDEX]])
"arm_sve.intr.psel"(%pn, %p1, %index) : (vector<[16]xi1>, vector<[2]xi1>, i32) -> vector<[16]xi1>
- // CHECK: call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv4i1(<vscale x 16 x i1> %[[PN]], <vscale x 4 x i1> %[[P2]], i32 %[[INDEX]])
+ // CHECK: call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1.nxv4i1(<vscale x 16 x i1> %[[PN]], <vscale x 4 x i1> %[[P2]], i32 %[[INDEX]])
"arm_sve.intr.psel"(%pn, %p2, %index) : (vector<[16]xi1>, vector<[4]xi1>, i32) -> vector<[16]xi1>
- // CHECK: call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv8i1(<vscale x 16 x i1> %[[PN]], <vscale x 8 x i1> %[[P3]], i32 %[[INDEX]])
+ // CHECK: call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1.nxv8i1(<vscale x 16 x i1> %[[PN]], <vscale x 8 x i1> %[[P3]], i32 %[[INDEX]])
"arm_sve.intr.psel"(%pn, %p3, %index) : (vector<[16]xi1>, vector<[8]xi1>, i32) -> vector<[16]xi1>
- // CHECK: call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1(<vscale x 16 x i1> %[[PN]], <vscale x 16 x i1> %[[P4]], i32 %[[INDEX]])
+ // CHECK: call <vscale x 16 x i1> @llvm.aarch64.sve.psel.nxv16i1.nxv16i1(<vscale x 16 x i1> %[[PN]], <vscale x 16 x i1> %[[P4]], i32 %[[INDEX]])
"arm_sve.intr.psel"(%pn, %p4, %index) : (vector<[16]xi1>, vector<[16]xi1>, i32) -> vector<[16]xi1>
llvm.return
}
More information about the Mlir-commits
mailing list