[llvm] [SPIR-V] Split wide vector intrinsics during legalization (PR #227287)
Arseniy Obolenskiy via llvm-commits
llvm-commits at lists.llvm.org
Tue Sep 29 05:13:48 PDT 2026
https://github.com/aobolensk created https://github.com/llvm/llvm-project/pull/227287
Fixes https://github.com/llvm/llvm-project/issues/225961
Requires https://github.com/llvm/llvm-project/pull/227286
>From 17e3215163684ba77ccd58a5f3dcd58550048b41 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 29 Sep 2026 14:09:40 +0200
Subject: [PATCH 1/2] [SPIR-V] Mark elementwise intrinsics
IntrTriviallyScalarizable
This mirrors the DirectX intrinsics
The SPIR-V legalizer will use it to split elementwise intrinsics with illegal vector widths
---
llvm/include/llvm/IR/IntrinsicsSPIRV.td | 84 ++++++++++++-------------
1 file changed, 42 insertions(+), 42 deletions(-)
diff --git a/llvm/include/llvm/IR/IntrinsicsSPIRV.td b/llvm/include/llvm/IR/IntrinsicsSPIRV.td
index 25b5a3c3854654..a2dad8b6a4ddb8 100644
--- a/llvm/include/llvm/IR/IntrinsicsSPIRV.td
+++ b/llvm/include/llvm/IR/IntrinsicsSPIRV.td
@@ -99,20 +99,20 @@ let TargetPrefix = "spv" in {
def int_spv_flattened_thread_id_in_group : Intrinsic<[llvm_i32_ty], [], [IntrNoMem, IntrWillReturn]>;
def int_spv_all : DefaultAttrsIntrinsic<[llvm_i1_ty], [llvm_any_ty], [IntrNoMem]>;
def int_spv_any : DefaultAttrsIntrinsic<[llvm_i1_ty], [llvm_any_ty], [IntrNoMem]>;
- def int_spv_degrees : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty], [IntrNoMem]>;
+ def int_spv_degrees : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty], [IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_distance : DefaultAttrsIntrinsic<[LLVMVectorElementType<0>], [llvm_anyfloat_ty, LLVMMatchType<0>], [IntrNoMem]>;
def int_spv_faceforward : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty, LLVMMatchType<0>, LLVMMatchType<0>], [IntrNoMem]>;
- def int_spv_frac : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty], [IntrNoMem]>;
+ def int_spv_frac : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty], [IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_isinf : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i1_ty>],
- [llvm_anyfloat_ty], [IntrNoMem]>;
+ [llvm_anyfloat_ty], [IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_isnan : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i1_ty>],
- [llvm_anyfloat_ty], [IntrNoMem]>;
+ [llvm_anyfloat_ty], [IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_isfinite : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i1_ty>],
- [llvm_anyfloat_ty], [IntrNoMem]>;
+ [llvm_anyfloat_ty], [IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_isnormal : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i1_ty>],
- [llvm_anyfloat_ty], [IntrNoMem]>;
+ [llvm_anyfloat_ty], [IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_lerp : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty, LLVMMatchType<0>,LLVMMatchType<0>],
- [IntrNoMem] >;
+ [IntrNoMem, IntrTriviallyScalarizable] >;
def int_spv_length : DefaultAttrsIntrinsic<[LLVMVectorElementType<0>], [llvm_anyfloat_ty], [IntrNoMem]>;
def int_spv_normalize : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty], [IntrNoMem]>;
def int_spv_reflect : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty, LLVMMatchType<0>], [IntrNoMem]>;
@@ -121,9 +121,9 @@ let TargetPrefix = "spv" in {
[llvm_anyfloat_ty, LLVMMatchType<0>,
llvm_anyfloat_ty],
[IntrNoMem]>;
-def int_spv_rsqrt : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty], [IntrNoMem]>;
- def int_spv_saturate : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem]>;
- def int_spv_smoothstep : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty, LLVMMatchType<0>, LLVMMatchType<0>], [IntrNoMem]>;
+def int_spv_rsqrt : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_saturate : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_smoothstep : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty, LLVMMatchType<0>, LLVMMatchType<0>], [IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_fdot :
DefaultAttrsIntrinsic<[LLVMVectorElementType<0>],
[llvm_anyfloat_ty, LLVMScalarOrSameVectorWidth<0, LLVMVectorElementType<0>>],
@@ -140,32 +140,32 @@ def int_spv_rsqrt : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty]
def int_spv_dot4add_u8packed : DefaultAttrsIntrinsic<[llvm_i32_ty], [llvm_i32_ty, llvm_i32_ty, llvm_i32_ty], [IntrNoMem]>;
def int_spv_subgroup_prefix_bit_count : DefaultAttrsIntrinsic<[llvm_i32_ty], [llvm_i1_ty], [IntrConvergent, IntrNoMem]>;
def int_spv_wave_active_countbits : DefaultAttrsIntrinsic<[llvm_i32_ty], [llvm_i1_ty], [IntrConvergent, IntrNoMem]>;
- def int_spv_wave_all_equal : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i1_ty>], [llvm_any_ty], [IntrConvergent, IntrNoMem]>;
+ def int_spv_wave_all_equal : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i1_ty>], [llvm_any_ty], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_wave_all : DefaultAttrsIntrinsic<[llvm_i1_ty], [llvm_i1_ty], [IntrConvergent, IntrNoMem]>;
def int_spv_wave_any : DefaultAttrsIntrinsic<[llvm_i1_ty], [llvm_i1_ty], [IntrConvergent, IntrNoMem]>;
- def int_spv_wave_reduce_or : DefaultAttrsIntrinsic<[llvm_anyint_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
- def int_spv_wave_reduce_xor : DefaultAttrsIntrinsic<[llvm_anyint_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
- def int_spv_wave_reduce_and : DefaultAttrsIntrinsic<[llvm_anyint_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
+ def int_spv_wave_reduce_or : DefaultAttrsIntrinsic<[llvm_anyint_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_wave_reduce_xor : DefaultAttrsIntrinsic<[llvm_anyint_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_wave_reduce_and : DefaultAttrsIntrinsic<[llvm_anyint_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_subgroup_ballot : ClangBuiltin<"__builtin_spirv_subgroup_ballot">,
DefaultAttrsIntrinsic<[llvm_v4i32_ty], [llvm_i1_ty], [IntrConvergent, IntrNoMem]>;
- def int_spv_wave_reduce_umax : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
- def int_spv_wave_reduce_max : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
- def int_spv_wave_reduce_min : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
- def int_spv_wave_reduce_umin : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
- def int_spv_wave_reduce_sum : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
- def int_spv_wave_product : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
+ def int_spv_wave_reduce_umax : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_wave_reduce_max : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_wave_reduce_min : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_wave_reduce_umin : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_wave_reduce_sum : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_wave_product : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_wave_is_first_lane : DefaultAttrsIntrinsic<[llvm_i1_ty], [], [IntrConvergent]>;
- def int_spv_wave_readlane : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>, llvm_i32_ty], [IntrConvergent, IntrNoMem]>;
- def int_spv_wave_readlane_first : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
+ def int_spv_wave_readlane : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>, llvm_i32_ty], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_wave_readlane_first : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_wave_get_lane_count
: DefaultAttrsIntrinsic<[llvm_i32_ty], [], [IntrConvergent]>;
- def int_spv_wave_prefix_sum : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
- def int_spv_wave_prefix_product : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
- def int_spv_quad_read_across_x : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
- def int_spv_quad_read_across_y : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
- def int_spv_quad_read_across_diagonal : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
- def int_spv_sign : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i32_ty>], [llvm_any_ty], [IntrNoMem]>;
- def int_spv_radians : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty], [IntrNoMem]>;
+ def int_spv_wave_prefix_sum : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_wave_prefix_product : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_quad_read_across_x : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_quad_read_across_y : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_quad_read_across_diagonal : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_sign : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i32_ty>], [llvm_any_ty], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_radians : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty], [IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_all_memory_barrier : DefaultAttrsIntrinsic<[], [], [IntrConvergent]>;
def int_spv_all_memory_barrier_with_group_sync : DefaultAttrsIntrinsic<[], [], [IntrConvergent]>;
def int_spv_device_memory_barrier : DefaultAttrsIntrinsic<[], [], [IntrConvergent]>;
@@ -174,16 +174,16 @@ def int_spv_rsqrt : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty]
def int_spv_group_memory_barrier_with_group_sync : ClangBuiltin<"__builtin_spirv_group_barrier">,
DefaultAttrsIntrinsic<[], [], [IntrConvergent]>;
def int_spv_discard : DefaultAttrsIntrinsic<[], [], []>;
- def int_spv_ddx : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem]>;
- def int_spv_ddy : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem]>;
- def int_spv_ddx_coarse : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem]>;
- def int_spv_ddy_coarse : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem]>;
- def int_spv_ddx_fine : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem]>;
- def int_spv_ddy_fine : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem]>;
- def int_spv_fwidth : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem]>;
- def int_spv_uclamp : DefaultAttrsIntrinsic<[llvm_anyint_ty], [LLVMMatchType<0>, LLVMMatchType<0>, LLVMMatchType<0>], [IntrNoMem]>;
- def int_spv_sclamp : DefaultAttrsIntrinsic<[llvm_anyint_ty], [LLVMMatchType<0>, LLVMMatchType<0>, LLVMMatchType<0>], [IntrNoMem]>;
- def int_spv_nclamp : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>, LLVMMatchType<0>, LLVMMatchType<0>], [IntrNoMem]>;
+ def int_spv_ddx : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_ddy : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_ddx_coarse : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_ddy_coarse : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_ddx_fine : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_ddy_fine : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_fwidth : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_uclamp : DefaultAttrsIntrinsic<[llvm_anyint_ty], [LLVMMatchType<0>, LLVMMatchType<0>, LLVMMatchType<0>], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_sclamp : DefaultAttrsIntrinsic<[llvm_anyint_ty], [LLVMMatchType<0>, LLVMMatchType<0>, LLVMMatchType<0>], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_nclamp : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [LLVMMatchType<0>, LLVMMatchType<0>, LLVMMatchType<0>], [IntrNoMem, IntrTriviallyScalarizable]>;
// Create resource handle given the binding information. Returns a
// type appropriate for the kind of resource given the set id, binding id,
@@ -217,9 +217,9 @@ def int_spv_rsqrt : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty]
[llvm_any_ty],
[IntrNoMem, IntrConvergent]>;
- def int_spv_firstbituhigh : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i32_ty>], [llvm_anyint_ty], [IntrNoMem]>;
- def int_spv_firstbitshigh : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i32_ty>], [llvm_anyint_ty], [IntrNoMem]>;
- def int_spv_firstbitlow : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i32_ty>], [llvm_anyint_ty], [IntrNoMem]>;
+ def int_spv_firstbituhigh : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i32_ty>], [llvm_anyint_ty], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_firstbitshigh : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i32_ty>], [llvm_anyint_ty], [IntrNoMem, IntrTriviallyScalarizable]>;
+ def int_spv_firstbitlow : DefaultAttrsIntrinsic<[LLVMScalarOrSameVectorWidth<0, llvm_i32_ty>], [llvm_anyint_ty], [IntrNoMem, IntrTriviallyScalarizable]>;
def int_spv_resource_updatecounter
: DefaultAttrsIntrinsic<[llvm_i32_ty], [llvm_any_ty, llvm_i8_ty],
>From 447be529830eeb45a9eaf8b2ff8cb65edc4ea856 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 29 Sep 2026 14:12:58 +0200
Subject: [PATCH 2/2] [SPIR-V] Split wide vector intrinsics during legalization
Fixes #225961
Requires #227286
---
llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp | 98 ++++++++++++++++++-
.../hlsl-intrinsics/WaveReadLaneFirst_mat.ll | 4 -
.../legalization/intrinsic-vector-split.ll | 46 +++++++++
3 files changed, 143 insertions(+), 5 deletions(-)
create mode 100644 llvm/test/CodeGen/SPIRV/legalization/intrinsic-vector-split.ll
diff --git a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
index b5d43a6074bc5e..dc8a5754e664ad 100644
--- a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
@@ -1051,6 +1051,102 @@ static bool legalizeSpvConstComposite(LegalizerHelper &Helper, MachineInstr &MI,
return true;
}
+static SmallVector<Register, 16> unmergeToScalars(Register Reg,
+ MachineIRBuilder &MIRBuilder,
+ SPIRVGlobalRegistry *GR) {
+ LLT Ty = MIRBuilder.getMRI()->getType(Reg);
+ if (!Ty.isVector())
+ return {Reg};
+ SPIRVTypeInst EltSpvTy =
+ GR->getScalarOrVectorComponentType(GR->getSPIRVTypeForVReg(Reg));
+ unsigned NumElts = Ty.getNumElements();
+ SmallVector<Register, 16> Elts;
+ for (unsigned I = 0; I < NumElts; ++I)
+ Elts.push_back(createVirtualRegister(EltSpvTy, GR, MIRBuilder));
+ MIRBuilder.buildUnmerge(Elts, Reg);
+ return Elts;
+}
+
+static Register buildVectorPart(ArrayRef<Register> Elts,
+ MachineIRBuilder &MIRBuilder,
+ SPIRVGlobalRegistry *GR) {
+ if (Elts.size() == 1)
+ return Elts[0];
+ SPIRVTypeInst PartSpvTy =
+ GR->getOrCreateSPIRVVectorType(GR->getSPIRVTypeForVReg(Elts[0]),
+ Elts.size(), MIRBuilder, /*EmitIR=*/true);
+ Register Part = createVirtualRegister(PartSpvTy, GR, MIRBuilder);
+ MIRBuilder.buildBuildVector(Part, Elts);
+ return Part;
+}
+
+// Split an elementwise intrinsic with an illegal vector width into intrinsics
+// on legal vector widths.
+static bool legalizeElementwiseIntrinsic(LegalizerHelper &Helper,
+ GIntrinsic &MI,
+ SPIRVGlobalRegistry *GR) {
+ MachineIRBuilder &MIRBuilder = Helper.MIRBuilder;
+ MachineRegisterInfo &MRI = *MIRBuilder.getMRI();
+ const SPIRVSubtarget &ST = MI.getMF()->getSubtarget<SPIRVSubtarget>();
+
+ if (!Intrinsic::isTriviallyScalarizable(MI.getIntrinsicID()))
+ return true;
+ Register DstReg = MI.getReg(0);
+ LLT DstTy = MRI.getType(DstReg);
+ if (!needsVectorLegalization(DstTy, ST))
+ return true;
+
+ unsigned NumElts = DstTy.getNumElements();
+ unsigned MaxVectorSize = ST.isShader() ? 4 : 16;
+ unsigned PartSize = NumElts > MaxVectorSize ? MaxVectorSize : 4;
+
+ SmallDenseMap<Register, SmallVector<Register, 16>, 4> OpElts;
+ for (const MachineOperand &MO : drop_begin(MI.explicit_uses())) {
+ if (!MO.isReg() || !MRI.getType(MO.getReg()).isVector())
+ continue;
+ auto [It, Inserted] = OpElts.try_emplace(MO.getReg());
+ if (Inserted)
+ It->second = unmergeToScalars(MO.getReg(), MIRBuilder, GR);
+ }
+
+ SPIRVTypeInst DstEltSpvTy =
+ GR->getScalarOrVectorComponentType(GR->getSPIRVTypeForVReg(DstReg));
+ SmallVector<Register, 16> DstElts;
+ for (unsigned Offset = 0; Offset < NumElts; Offset += PartSize) {
+ unsigned Size = std::min(PartSize, NumElts - Offset);
+ SPIRVTypeInst PartSpvTy =
+ Size == 1 ? DstEltSpvTy
+ : GR->getOrCreateSPIRVVectorType(DstEltSpvTy, Size,
+ MIRBuilder, /*EmitIR=*/true);
+ SmallDenseMap<Register, Register, 4> PartRegs;
+ SmallVector<MachineOperand> PartOps;
+ for (const MachineOperand &MO : drop_begin(MI.explicit_uses())) {
+ auto EltsIt = MO.isReg() ? OpElts.find(MO.getReg()) : OpElts.end();
+ if (EltsIt == OpElts.end()) {
+ PartOps.push_back(MO);
+ continue;
+ }
+ auto [It, Inserted] = PartRegs.try_emplace(MO.getReg());
+ if (Inserted)
+ It->second = buildVectorPart(
+ ArrayRef(EltsIt->second).slice(Offset, Size), MIRBuilder, GR);
+ PartOps.push_back(MachineOperand::CreateReg(It->second, /*isDef=*/false));
+ }
+ Register PartDst = createVirtualRegister(PartSpvTy, GR, MIRBuilder);
+ auto Part = MIRBuilder.buildIntrinsic(
+ MI.getIntrinsicID(), ArrayRef<Register>{PartDst}, MI.hasSideEffects(),
+ MI.isConvergent());
+ for (const MachineOperand &MO : PartOps)
+ Part.add(MO);
+ Part->setFlags(MI.getFlags());
+ append_range(DstElts, unmergeToScalars(PartDst, MIRBuilder, GR));
+ }
+
+ MIRBuilder.buildBuildVector(DstReg, DstElts);
+ MI.eraseFromParent();
+ return true;
+}
+
bool SPIRVLegalizerInfo::legalizeIntrinsic(LegalizerHelper &Helper,
MachineInstr &MI) const {
LLVM_DEBUG(dbgs() << "legalizeIntrinsic: " << MI);
@@ -1065,7 +1161,7 @@ bool SPIRVLegalizerInfo::legalizeIntrinsic(LegalizerHelper &Helper,
case Intrinsic::spv_const_composite:
return legalizeSpvConstComposite(Helper, MI, GR);
}
- return true;
+ return legalizeElementwiseIntrinsic(Helper, cast<GIntrinsic>(MI), GR);
}
bool SPIRVLegalizerInfo::legalizeBitcast(LegalizerHelper &Helper,
diff --git a/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst_mat.ll b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst_mat.ll
index 02517fb6571714..e1aa0ecfd7aed6 100644
--- a/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst_mat.ll
+++ b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst_mat.ll
@@ -1,7 +1,3 @@
-; XFAIL: *
-; TODO: Support matrix legalization for SPIR-V target intrinsics.
-; https://github.com/llvm/llvm-project/issues/225961
-;
; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv1.6-unknown-vulkan1.3-compute %s -o - | FileCheck %s
; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv1.6-unknown-vulkan1.3-compute %s -o - -filetype=obj | spirv-val --target-env vulkan1.3 %}
diff --git a/llvm/test/CodeGen/SPIRV/legalization/intrinsic-vector-split.ll b/llvm/test/CodeGen/SPIRV/legalization/intrinsic-vector-split.ll
new file mode 100644
index 00000000000000..7559d253ef3bef
--- /dev/null
+++ b/llvm/test/CodeGen/SPIRV/legalization/intrinsic-vector-split.ll
@@ -0,0 +1,46 @@
+; RUN: llc -O0 -verify-machineinstrs -mtriple=spirv1.6-unknown-vulkan1.3-compute %s -o - | FileCheck %s
+; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv1.6-unknown-vulkan1.3-compute %s -o - -filetype=obj | spirv-val --target-env vulkan1.3 %}
+
+; CHECK-DAG: %[[#glsl:]] = OpExtInstImport "GLSL.std.450"
+; CHECK-DAG: %[[#uint:]] = OpTypeInt 32 0
+; CHECK-DAG: %[[#f32:]] = OpTypeFloat 32
+; CHECK-DAG: %[[#bool:]] = OpTypeBool
+; CHECK-DAG: %[[#v2f32:]] = OpTypeVector %[[#f32]] 2
+; CHECK-DAG: %[[#v4f32:]] = OpTypeVector %[[#f32]] 4
+; CHECK-DAG: %[[#v2bool:]] = OpTypeVector %[[#bool]] 2
+; CHECK-DAG: %[[#v4bool:]] = OpTypeVector %[[#bool]] 4
+; CHECK-DAG: %[[#scope:]] = OpConstant %[[#uint]] 3
+
+ at f = internal addrspace(10) global [6 x float] zeroinitializer
+ at i = internal addrspace(10) global [6 x i32] zeroinitializer
+
+define void @main() #0 {
+entry:
+ %v = load <6 x float>, ptr addrspace(10) @f
+
+; The scalar lane index is reused by every part.
+; CHECK: %[[#rl0:]] = OpGroupNonUniformShuffle %[[#v4f32]] %[[#scope]] %[[#]] %[[#scope]]
+; CHECK: %[[#rl1:]] = OpGroupNonUniformShuffle %[[#v2f32]] %[[#scope]] %[[#]] %[[#scope]]
+ %rl = call <6 x float> @llvm.spv.wave.readlane.v6f32(<6 x float> %v, i32 3)
+
+; CHECK: %[[#cl0:]] = OpExtInst %[[#v4f32]] %[[#glsl]] NClamp %[[#rl0]] %[[#]] %[[#]]
+; CHECK: %[[#cl1:]] = OpExtInst %[[#v2f32]] %[[#glsl]] NClamp %[[#rl1]] %[[#]] %[[#]]
+ %cl = call <6 x float> @llvm.spv.nclamp.v6f32(<6 x float> %rl, <6 x float> %v, <6 x float> %v)
+ store <6 x float> %cl, ptr addrspace(10) @f
+
+; CHECK: OpIsInf %[[#v4bool]] %[[#cl0]]
+; CHECK: OpIsInf %[[#v2bool]] %[[#cl1]]
+ %inf = call <6 x i1> @llvm.spv.isinf.v6f32(<6 x float> %cl)
+ %ext = zext <6 x i1> %inf to <6 x i32>
+ store <6 x i32> %ext, ptr addrspace(10) @i
+
+; A single trailing element becomes a scalar call.
+; CHECK: OpExtInst %[[#v4f32]] %[[#glsl]] Fract
+; CHECK: OpExtInst %[[#f32]] %[[#glsl]] Fract
+ %v5 = load <5 x float>, ptr addrspace(10) @f
+ %fr = call <5 x float> @llvm.spv.frac.v5f32(<5 x float> %v5)
+ store <5 x float> %fr, ptr addrspace(10) @f
+ ret void
+}
+
+attributes #0 = { "hlsl.numthreads"="1,1,1" "hlsl.shader"="compute" }
More information about the llvm-commits
mailing list