[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