[clang] [llvm] Add WaveReadLaneFirst HLSL function (PR #220373)
Joshua Batista via llvm-commits
llvm-commits at lists.llvm.org
Mon Sep 21 17:47:41 PDT 2026
https://github.com/bob80905 updated https://github.com/llvm/llvm-project/pull/220373
>From 041d800aed222bb150ef4d321529edc567c6896e Mon Sep 17 00:00:00 2001
From: Joshua Batista <jbatista at microsoft.com>
Date: Tue, 1 Sep 2026 13:27:20 -0700
Subject: [PATCH 1/7] first attempt
---
clang/include/clang/Basic/Builtins.td | 6 +
clang/include/clang/Basic/HLSLIntrinsics.td | 14 +++
clang/lib/CodeGen/CGHLSLBuiltins.cpp | 6 +
clang/lib/CodeGen/CGHLSLRuntime.h | 1 +
clang/lib/Sema/SemaHLSL.cpp | 10 ++
.../builtins/WaveReadLaneFirst.hlsl | 109 ++++++++++++++++++
.../BuiltIns/WaveReadLaneFirst-errors.hlsl | 18 +++
.../SemaHLSL/WaveBuiltinAvailability.hlsl | 8 ++
llvm/include/llvm/IR/IntrinsicsDirectX.td | 1 +
llvm/include/llvm/IR/IntrinsicsSPIRV.td | 1 +
llvm/lib/Target/DirectX/DXIL.td | 10 ++
llvm/lib/Target/DirectX/DXILShaderFlags.cpp | 1 +
.../Target/SPIRV/SPIRVInstructionSelector.cpp | 3 +
.../CodeGen/DirectX/ShaderFlags/wave-ops.ll | 7 ++
.../test/CodeGen/DirectX/WaveReadLaneFirst.ll | 83 +++++++++++++
.../hlsl-intrinsics/WaveReadLaneFirst.ll | 55 +++++++++
16 files changed, 333 insertions(+)
create mode 100644 clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl
create mode 100644 clang/test/SemaHLSL/BuiltIns/WaveReadLaneFirst-errors.hlsl
create mode 100644 llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll
create mode 100644 llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
diff --git a/clang/include/clang/Basic/Builtins.td b/clang/include/clang/Basic/Builtins.td
index 49fe879c6add15..73b73623d81f64 100644
--- a/clang/include/clang/Basic/Builtins.td
+++ b/clang/include/clang/Basic/Builtins.td
@@ -5629,6 +5629,12 @@ def HLSLWaveReadLaneAt : LangBuiltin<"HLSL_LANG"> {
let Prototype = "void(...)";
}
+def HLSLWaveReadLaneFirst : LangBuiltin<"HLSL_LANG"> {
+ let Spellings = ["__builtin_hlsl_wave_read_lane_first"];
+ let Attributes = [NoThrow, Const];
+ let Prototype = "void(...)";
+}
+
def HLSLWaveGetLaneCount : LangBuiltin<"HLSL_LANG"> {
let Spellings = ["__builtin_hlsl_wave_get_lane_count"];
let Attributes = [NoThrow, Const];
diff --git a/clang/include/clang/Basic/HLSLIntrinsics.td b/clang/include/clang/Basic/HLSLIntrinsics.td
index 4ef818edb32595..d4f6fa8c5357d9 100644
--- a/clang/include/clang/Basic/HLSLIntrinsics.td
+++ b/clang/include/clang/Basic/HLSLIntrinsics.td
@@ -1936,3 +1936,17 @@ the specified wave.
let Availability = SM6_0;
let VaryingMatDims = [];
}
+
+// Reads the value from the first active lane in the wave.
+def hlsl_wave_read_lane_first :
+ HLSLOneArgBuiltin<"WaveReadLaneFirst",
+ "__builtin_hlsl_wave_read_lane_first"> {
+ let Doc = [{
+\brief Returns the value from the active lane with the smallest index.
+\param Val The value to read.
+}];
+ let VaryingTypes = AllTypesWithBool;
+ let IsConvergent = 1;
+ let Availability = SM6_0;
+ let VaryingMatDims = [];
+}
diff --git a/clang/lib/CodeGen/CGHLSLBuiltins.cpp b/clang/lib/CodeGen/CGHLSLBuiltins.cpp
index 062faadcdcab27..75716021606733 100644
--- a/clang/lib/CodeGen/CGHLSLBuiltins.cpp
+++ b/clang/lib/CodeGen/CGHLSLBuiltins.cpp
@@ -1580,6 +1580,12 @@ Value *CodeGenFunction::EmitHLSLBuiltinExpr(unsigned BuiltinID,
{OpExpr->getType()}, ArrayRef{OpExpr, OpIndex},
"hlsl.wave.readlane");
}
+ case Builtin::BI__builtin_hlsl_wave_read_lane_first: {
+ Value *OpExpr = EmitScalarExpr(E->getArg(0));
+ return EmitIntrinsicCall(
+ CGM.getHLSLRuntime().getWaveReadLaneFirstIntrinsic(),
+ {OpExpr->getType()}, ArrayRef{OpExpr}, "hlsl.wave.readlane.first");
+ }
case Builtin::BI__builtin_hlsl_wave_prefix_sum: {
Value *OpExpr = EmitScalarExpr(E->getArg(0));
Intrinsic::ID IID = getWavePrefixSumIntrinsic(
diff --git a/clang/lib/CodeGen/CGHLSLRuntime.h b/clang/lib/CodeGen/CGHLSLRuntime.h
index 381653e8f83456..a599750c8c26b1 100644
--- a/clang/lib/CodeGen/CGHLSLRuntime.h
+++ b/clang/lib/CodeGen/CGHLSLRuntime.h
@@ -154,6 +154,7 @@ class CGHLSLRuntime {
GENERATE_HLSL_INTRINSIC_FUNCTION(WaveIsFirstLane, wave_is_first_lane)
GENERATE_HLSL_INTRINSIC_FUNCTION(WaveGetLaneCount, wave_get_lane_count)
GENERATE_HLSL_INTRINSIC_FUNCTION(WaveReadLaneAt, wave_readlane)
+ GENERATE_HLSL_INTRINSIC_FUNCTION(WaveReadLaneFirst, wave_readlane_first)
GENERATE_HLSL_INTRINSIC_FUNCTION(QuadReadAcrossX, quad_read_across_x)
GENERATE_HLSL_INTRINSIC_FUNCTION(QuadReadAcrossY, quad_read_across_y)
GENERATE_HLSL_INTRINSIC_FUNCTION(QuadReadAcrossDiagonal,
diff --git a/clang/lib/Sema/SemaHLSL.cpp b/clang/lib/Sema/SemaHLSL.cpp
index 06828b9ec7fc0a..7309dd79f065a0 100644
--- a/clang/lib/Sema/SemaHLSL.cpp
+++ b/clang/lib/Sema/SemaHLSL.cpp
@@ -4829,6 +4829,16 @@ bool SemaHLSL::CheckBuiltinFunctionCall(unsigned BuiltinID, CallExpr *TheCall) {
TheCall->setType(ArgTyExpr);
break;
}
+ case Builtin::BI__builtin_hlsl_wave_read_lane_first: {
+ if (SemaRef.checkArgCount(TheCall, 1))
+ return true;
+
+ if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
+ return true;
+
+ TheCall->setType(TheCall->getArg(0)->getType());
+ break;
+ }
case Builtin::BI__builtin_hlsl_wave_get_lane_index: {
if (SemaRef.checkArgCount(TheCall, 0))
return true;
diff --git a/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl b/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl
new file mode 100644
index 00000000000000..79ddba9d9cc22e
--- /dev/null
+++ b/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl
@@ -0,0 +1,109 @@
+// RUN: %clang_cc1 -std=hlsl2021 -finclude-default-header -fnative-half-type -fnative-int16-type -triple \
+// RUN: dxil-pc-shadermodel6.3-library %s -emit-llvm -disable-llvm-passes -o - | \
+// RUN: FileCheck %s --check-prefixes=CHECK,CHECK-DXIL
+// RUN: %clang_cc1 -std=hlsl2021 -finclude-default-header -fnative-half-type -fnative-int16-type -triple \
+// RUN: spirv-pc-vulkan-library %s -emit-llvm -disable-llvm-passes -o - | \
+// RUN: FileCheck %s --check-prefixes=CHECK,CHECK-SPIRV
+
+// CHECK-LABEL: test_int
+int test_int(int expr) {
+ // CHECK-SPIRV: %[[#entry_tok0:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok0]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i32([[TY]] %[[#]])
+ // CHECK: ret [[TY]] %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+
+// CHECK-DXIL: declare [[TY]] @llvm.dx.wave.readlane.first.i32([[TY]]) #[[#attr:]]
+// CHECK-SPIRV: declare [[TY]] @llvm.spv.wave.readlane.first.i32([[TY]]) #[[#attr:]]
+
+// CHECK-LABEL: test_uint
+uint test_uint(uint expr) {
+ // CHECK-SPIRV: %[[#entry_tok0:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok0]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i32([[TY]] %[[#]])
+ // CHECK: ret [[TY]] %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+
+// CHECK-LABEL: test_int64_t
+int64_t test_int64_t(int64_t expr) {
+ // CHECK-SPIRV: %[[#entry_tok1:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok1]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i64([[TY]] %[[#]])
+ // CHECK: ret [[TY]] %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+
+// CHECK-DXIL: declare [[TY]] @llvm.dx.wave.readlane.first.i64([[TY]]) #[[#attr:]]
+// CHECK-SPIRV: declare [[TY]] @llvm.spv.wave.readlane.first.i64([[TY]]) #[[#attr:]]
+
+// CHECK-LABEL: test_uint64_t
+uint64_t test_uint64_t(uint64_t expr) {
+ // CHECK-SPIRV: %[[#entry_tok1:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok1]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i64([[TY]] %[[#]])
+ // CHECK: ret [[TY]] %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+
+#ifdef __HLSL_ENABLE_16_BIT
+// CHECK-LABEL: test_int16
+int16_t test_int16(int16_t expr) {
+ // CHECK-SPIRV: %[[#entry_tok2:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok2]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i16([[TY]] %[[#]])
+ // CHECK: ret [[TY]] %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+
+// CHECK-DXIL: declare [[TY]] @llvm.dx.wave.readlane.first.i16([[TY]]) #[[#attr:]]
+// CHECK-SPIRV: declare [[TY]] @llvm.spv.wave.readlane.first.i16([[TY]]) #[[#attr:]]
+
+// CHECK-LABEL: test_uint16
+uint16_t test_uint16(uint16_t expr) {
+ // CHECK-SPIRV: %[[#entry_tok2:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok2]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i16([[TY]] %[[#]])
+ // CHECK: ret [[TY]] %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+#endif
+
+// CHECK-LABEL: test_bool
+bool test_bool(bool expr) {
+ // CHECK-SPIRV: %[[#entry_tok3:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call i1 @llvm.spv.wave.readlane.first.i1(i1 %{{[a-zA-Z0-9]+}}) [ "convergencectrl"(token %[[#entry_tok3]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call i1 @llvm.dx.wave.readlane.first.i1(i1 %{{[a-zA-Z0-9]+}})
+ // CHECK: ret i1 %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+
+// CHECK-LABEL: test_half
+half test_half(half expr) {
+ // CHECK-SPIRV: %[[#entry_tok4:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.f16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok4]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.f16([[TY]] %[[#]])
+ // CHECK: ret [[TY]] %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+
+// CHECK-LABEL: test_double
+double test_double(double expr) {
+ // CHECK-SPIRV: %[[#entry_tok5:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.f64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok5]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.f64([[TY]] %[[#]])
+ // CHECK: ret [[TY]] %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+
+// CHECK-LABEL: test_floatv4
+float4 test_floatv4(float4 expr) {
+ // CHECK-SPIRV: %[[#entry_tok6:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.v4f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok6]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.v4f32([[TY]] %[[#]])
+ // CHECK: ret [[TY]] %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+
+// CHECK: attributes #[[#attr]] = {{{.*}} convergent {{.*}}}
diff --git a/clang/test/SemaHLSL/BuiltIns/WaveReadLaneFirst-errors.hlsl b/clang/test/SemaHLSL/BuiltIns/WaveReadLaneFirst-errors.hlsl
new file mode 100644
index 00000000000000..b042362cb79aff
--- /dev/null
+++ b/clang/test/SemaHLSL/BuiltIns/WaveReadLaneFirst-errors.hlsl
@@ -0,0 +1,18 @@
+// RUN: %clang_cc1 -finclude-default-header -triple dxil-pc-shadermodel6.6-library %s -emit-llvm-only -disable-llvm-passes -verify
+
+bool test_too_few_arg() {
+ return __builtin_hlsl_wave_read_lane_first();
+ // expected-error at -1 {{too few arguments to function call, expected 1, have 0}}
+}
+
+float2 test_too_many_arg(float2 p0) {
+ return __builtin_hlsl_wave_read_lane_first(p0, p0);
+ // expected-error at -1 {{too many arguments to function call, expected 1, have 2}}
+}
+
+struct S { float f; };
+
+S test_expr_struct_type_check(S p0) {
+ return __builtin_hlsl_wave_read_lane_first(p0);
+ // expected-error at -1 {{invalid operand of type 'S' where a scalar or vector is required}}
+}
diff --git a/clang/test/SemaHLSL/WaveBuiltinAvailability.hlsl b/clang/test/SemaHLSL/WaveBuiltinAvailability.hlsl
index 5741b81832ad05..7501d895e81634 100644
--- a/clang/test/SemaHLSL/WaveBuiltinAvailability.hlsl
+++ b/clang/test/SemaHLSL/WaveBuiltinAvailability.hlsl
@@ -36,6 +36,10 @@ void foo() {
// expected-note at hlsl/hlsl_alias_intrinsics_gen.inc:* {{'WaveReadLaneAt' has been marked as being introduced in Shader Model 6.0 here, but the deployment target is Shader Model 5.0}}
float g = hlsl::WaveReadLaneAt(1.0f, 0u); // #WaveReadLaneAt
+ // expected-error@#WaveReadLaneFirst {{'WaveReadLaneFirst' is only available on Shader Model 6.0 or newer}}
+ // expected-note at hlsl/hlsl_alias_intrinsics_gen.inc:* {{'WaveReadLaneFirst' has been marked as being introduced in Shader Model 6.0 here, but the deployment target is Shader Model 5.0}}
+ float first = hlsl::WaveReadLaneFirst(1.0f); // #WaveReadLaneFirst
+
// Test that half overloads (which map to float without native half) also
// have the correct SM 6.0 availability via the _HLSL_16BIT_AVAILABILITY
// fallback path.
@@ -46,4 +50,8 @@ void foo() {
// expected-error@#WaveReadLaneAt_half {{'WaveReadLaneAt' is only available on Shader Model 6.0 or newer}}
// expected-note at hlsl/hlsl_alias_intrinsics_gen.inc:* {{'WaveReadLaneAt' has been marked as being introduced in Shader Model 6.0 here, but the deployment target is Shader Model 5.0}}
half i = hlsl::WaveReadLaneAt((half)1.0, 0u); // #WaveReadLaneAt_half
+
+ // expected-error@#WaveReadLaneFirst_half {{'WaveReadLaneFirst' is only available on Shader Model 6.0 or newer}}
+ // expected-note at hlsl/hlsl_alias_intrinsics_gen.inc:* {{'WaveReadLaneFirst' has been marked as being introduced in Shader Model 6.0 here, but the deployment target is Shader Model 5.0}}
+ half j = hlsl::WaveReadLaneFirst((half)1.0); // #WaveReadLaneFirst_half
}
diff --git a/llvm/include/llvm/IR/IntrinsicsDirectX.td b/llvm/include/llvm/IR/IntrinsicsDirectX.td
index a8927f83ee2f8a..e5b5aa6648e372 100644
--- a/llvm/include/llvm/IR/IntrinsicsDirectX.td
+++ b/llvm/include/llvm/IR/IntrinsicsDirectX.td
@@ -288,6 +288,7 @@ def int_dx_wave_product : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>
def int_dx_wave_uproduct : DefaultAttrsIntrinsic<[llvm_anyint_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
def int_dx_wave_is_first_lane : DefaultAttrsIntrinsic<[llvm_i1_ty], [], [IntrConvergent]>;
def int_dx_wave_readlane : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>, llvm_i32_ty], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
+def int_dx_wave_readlane_first : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
def int_dx_wave_get_lane_count
: DefaultAttrsIntrinsic<[llvm_i32_ty], [], [IntrConvergent]>;
def int_dx_wave_prefix_sum : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem, IntrTriviallyScalarizable]>;
diff --git a/llvm/include/llvm/IR/IntrinsicsSPIRV.td b/llvm/include/llvm/IR/IntrinsicsSPIRV.td
index 86b49a8ee446a2..d449ecb303aa67 100644
--- a/llvm/include/llvm/IR/IntrinsicsSPIRV.td
+++ b/llvm/include/llvm/IR/IntrinsicsSPIRV.td
@@ -156,6 +156,7 @@ def int_spv_rsqrt : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty]
def int_spv_wave_product : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
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_get_lane_count
: DefaultAttrsIntrinsic<[llvm_i32_ty], [], [IntrConvergent]>;
def int_spv_wave_prefix_sum : DefaultAttrsIntrinsic<[llvm_any_ty], [LLVMMatchType<0>], [IntrConvergent, IntrNoMem]>;
diff --git a/llvm/lib/Target/DirectX/DXIL.td b/llvm/lib/Target/DirectX/DXIL.td
index 4beafd0c619b0e..0eee9457e869fa 100644
--- a/llvm/lib/Target/DirectX/DXIL.td
+++ b/llvm/lib/Target/DirectX/DXIL.td
@@ -1272,6 +1272,16 @@ def WaveReadLaneAt : DXILOp<117, waveReadLaneAt> {
let stages = [Stages<DXIL1_0, [all_stages]>];
}
+def WaveReadLaneFirst : DXILOp<118, waveReadLaneFirst> {
+ let Doc = "returns the value from the first active lane";
+ let intrinsics = [IntrinSelect<int_dx_wave_readlane_first>];
+ let arguments = [OverloadTy];
+ let result = OverloadTy;
+ let overloads = [Overloads<
+ DXIL1_0, [HalfTy, FloatTy, DoubleTy, Int1Ty, Int16Ty, Int32Ty, Int64Ty]>];
+ let stages = [Stages<DXIL1_0, [all_stages]>];
+}
+
def WaveActiveOp : DXILOp<119, waveActiveOp> {
let Doc = "returns the result of the operation across waves";
let intrinsics = [
diff --git a/llvm/lib/Target/DirectX/DXILShaderFlags.cpp b/llvm/lib/Target/DirectX/DXILShaderFlags.cpp
index e404eb81097694..6b8cbb850009b8 100644
--- a/llvm/lib/Target/DirectX/DXILShaderFlags.cpp
+++ b/llvm/lib/Target/DirectX/DXILShaderFlags.cpp
@@ -89,6 +89,7 @@ static bool checkWaveOps(Intrinsic::ID IID) {
case Intrinsic::dx_wave_all_equal:
case Intrinsic::dx_wave_all:
case Intrinsic::dx_wave_readlane:
+ case Intrinsic::dx_wave_readlane_first:
case Intrinsic::dx_wave_active_countbits:
case Intrinsic::dx_wave_ballot:
case Intrinsic::dx_wave_prefix_bit_count:
diff --git a/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp b/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
index 891a8d9da12cfd..8275e62802d1ad 100644
--- a/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
@@ -5772,6 +5772,9 @@ bool SPIRVInstructionSelector::selectIntrinsic(Register ResVReg,
case Intrinsic::spv_wave_readlane:
return selectWaveOpInst(ResVReg, ResType, I,
SPIRV::OpGroupNonUniformShuffle);
+ case Intrinsic::spv_wave_readlane_first:
+ return selectWaveOpInst(ResVReg, ResType, I,
+ SPIRV::OpGroupNonUniformBroadcastFirst);
case Intrinsic::spv_wave_prefix_sum:
return selectWaveExclusiveScanSum(ResVReg, ResType, I);
case Intrinsic::spv_wave_prefix_product:
diff --git a/llvm/test/CodeGen/DirectX/ShaderFlags/wave-ops.ll b/llvm/test/CodeGen/DirectX/ShaderFlags/wave-ops.ll
index 77478add84caf5..37a991e1c1fa8e 100644
--- a/llvm/test/CodeGen/DirectX/ShaderFlags/wave-ops.ll
+++ b/llvm/test/CodeGen/DirectX/ShaderFlags/wave-ops.ll
@@ -84,6 +84,13 @@ entry:
ret i1 %ret
}
+define noundef i1 @wave_readlane_first(i1 %x) {
+entry:
+ ; CHECK: Function wave_readlane_first : [[WAVE_FLAG]]
+ %ret = call i1 @llvm.dx.wave.readlane.first.i1(i1 %x)
+ ret i1 %ret
+}
+
define noundef i32 @wave_reduce_sum(i32 noundef %x) {
entry:
; CHECK: Function wave_reduce_sum : [[WAVE_FLAG]]
diff --git a/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll b/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll
new file mode 100644
index 00000000000000..336e4ba90adea9
--- /dev/null
+++ b/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll
@@ -0,0 +1,83 @@
+; RUN: opt -S -scalarizer -dxil-op-lower -mtriple=dxil-pc-shadermodel6.3-compute %s | FileCheck %s
+
+; Test that WaveReadLaneFirst maps down to the DirectX op.
+
+define noundef half @wave_readlane_first_half(half noundef %expr) {
+entry:
+; CHECK: call half @dx.op.waveReadLaneFirst.f16(i32 118, half %expr)
+ %ret = call half @llvm.dx.wave.readlane.first.f16(half %expr)
+ ret half %ret
+}
+
+define noundef float @wave_readlane_first_float(float noundef %expr) {
+entry:
+; CHECK: call float @dx.op.waveReadLaneFirst.f32(i32 118, float %expr)
+ %ret = call float @llvm.dx.wave.readlane.first.f32(float %expr)
+ ret float %ret
+}
+
+define noundef double @wave_readlane_first_double(double noundef %expr) {
+entry:
+; CHECK: call double @dx.op.waveReadLaneFirst.f64(i32 118, double %expr)
+ %ret = call double @llvm.dx.wave.readlane.first.f64(double %expr)
+ ret double %ret
+}
+
+define noundef i1 @wave_readlane_first_i1(i1 noundef %expr) {
+entry:
+; CHECK: call i1 @dx.op.waveReadLaneFirst.i1(i32 118, i1 %expr)
+ %ret = call i1 @llvm.dx.wave.readlane.first.i1(i1 %expr)
+ ret i1 %ret
+}
+
+define noundef i16 @wave_readlane_first_i16(i16 noundef %expr) {
+entry:
+; CHECK: call i16 @dx.op.waveReadLaneFirst.i16(i32 118, i16 %expr)
+ %ret = call i16 @llvm.dx.wave.readlane.first.i16(i16 %expr)
+ ret i16 %ret
+}
+
+define noundef i32 @wave_readlane_first_i32(i32 noundef %expr) {
+entry:
+; CHECK: call i32 @dx.op.waveReadLaneFirst.i32(i32 118, i32 %expr)
+ %ret = call i32 @llvm.dx.wave.readlane.first.i32(i32 %expr)
+ ret i32 %ret
+}
+
+define noundef i64 @wave_readlane_first_i64(i64 noundef %expr) {
+entry:
+; CHECK: call i64 @dx.op.waveReadLaneFirst.i64(i32 118, i64 %expr)
+ %ret = call i64 @llvm.dx.wave.readlane.first.i64(i64 %expr)
+ ret i64 %ret
+}
+
+define noundef <2 x half> @wave_readlane_first_v2half(
+ <2 x half> noundef %expr) {
+entry:
+; CHECK: call half @dx.op.waveReadLaneFirst.f16(i32 118, half %expr.i0)
+; CHECK: call half @dx.op.waveReadLaneFirst.f16(i32 118, half %expr.i1)
+ %ret = call <2 x half> @llvm.dx.wave.readlane.first.v2f16(
+ <2 x half> %expr)
+ ret <2 x half> %ret
+}
+
+define noundef <3 x i32> @wave_readlane_first_v3i32(
+ <3 x i32> noundef %expr) {
+entry:
+; CHECK: call i32 @dx.op.waveReadLaneFirst.i32(i32 118, i32 %expr.i0)
+; CHECK: call i32 @dx.op.waveReadLaneFirst.i32(i32 118, i32 %expr.i1)
+; CHECK: call i32 @dx.op.waveReadLaneFirst.i32(i32 118, i32 %expr.i2)
+ %ret = call <3 x i32> @llvm.dx.wave.readlane.first.v3i32(
+ <3 x i32> %expr)
+ ret <3 x i32> %ret
+}
+
+declare half @llvm.dx.wave.readlane.first.f16(half)
+declare float @llvm.dx.wave.readlane.first.f32(float)
+declare double @llvm.dx.wave.readlane.first.f64(double)
+declare i1 @llvm.dx.wave.readlane.first.i1(i1)
+declare i16 @llvm.dx.wave.readlane.first.i16(i16)
+declare i32 @llvm.dx.wave.readlane.first.i32(i32)
+declare i64 @llvm.dx.wave.readlane.first.i64(i64)
+declare <2 x half> @llvm.dx.wave.readlane.first.v2f16(<2 x half>)
+declare <3 x i32> @llvm.dx.wave.readlane.first.v3i32(<3 x i32>)
diff --git a/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
new file mode 100644
index 00000000000000..67651fee276973
--- /dev/null
+++ b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
@@ -0,0 +1,55 @@
+; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv1.5-vulkan-unknown %s -o - | FileCheck %s
+; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv1.5-vulkan-unknown %s -o - -filetype=obj | spirv-val %}
+
+; Test WaveReadLaneFirst lowering for scalar and vector types.
+
+; CHECK: Capability Shader
+; CHECK: Capability GroupNonUniformBallot
+
+; CHECK-DAG: %[[#uint:]] = OpTypeInt 32 0
+; CHECK-DAG: %[[#f32:]] = OpTypeFloat 32
+; CHECK-DAG: %[[#v4_float:]] = OpTypeVector %[[#f32]] 4
+; CHECK-DAG: %[[#bool:]] = OpTypeBool
+; CHECK-DAG: %[[#scope:]] = OpConstant %[[#uint]] 3
+
+; CHECK-LABEL: Begin function test_float
+; CHECK: %[[#fexpr:]] = OpFunctionParameter %[[#f32]]
+define float @test_float(float %fexpr) {
+entry:
+; CHECK: %[[#]] = OpGroupNonUniformBroadcastFirst %[[#f32]] %[[#scope]] %[[#fexpr]]
+ %0 = call float @llvm.spv.wave.readlane.first.f32(float %fexpr)
+ ret float %0
+}
+
+; CHECK-LABEL: Begin function test_int
+; CHECK: %[[#iexpr:]] = OpFunctionParameter %[[#uint]]
+define i32 @test_int(i32 %iexpr) {
+entry:
+; CHECK: %[[#]] = OpGroupNonUniformBroadcastFirst %[[#uint]] %[[#scope]] %[[#iexpr]]
+ %0 = call i32 @llvm.spv.wave.readlane.first.i32(i32 %iexpr)
+ ret i32 %0
+}
+
+; CHECK-LABEL: Begin function test_bool
+; CHECK: %[[#bexpr:]] = OpFunctionParameter %[[#bool]]
+define i1 @test_bool(i1 %bexpr) {
+entry:
+; CHECK: %[[#]] = OpGroupNonUniformBroadcastFirst %[[#bool]] %[[#scope]] %[[#bexpr]]
+ %0 = call i1 @llvm.spv.wave.readlane.first.i1(i1 %bexpr)
+ ret i1 %0
+}
+
+; CHECK-LABEL: Begin function test_vfloat
+; CHECK: %[[#vfexpr:]] = OpFunctionParameter %[[#v4_float]]
+define <4 x float> @test_vfloat(<4 x float> %vfexpr) {
+entry:
+; CHECK: %[[#]] = OpGroupNonUniformBroadcastFirst %[[#v4_float]] %[[#scope]] %[[#vfexpr]]
+ %0 = call <4 x float> @llvm.spv.wave.readlane.first.v4f32(
+ <4 x float> %vfexpr)
+ ret <4 x float> %0
+}
+
+declare float @llvm.spv.wave.readlane.first.f32(float)
+declare i32 @llvm.spv.wave.readlane.first.i32(i32)
+declare i1 @llvm.spv.wave.readlane.first.i1(i1)
+declare <4 x float> @llvm.spv.wave.readlane.first.v4f32(<4 x float>)
>From a848827593579eb5d0a11dfd05078bfdd7ffe6c7 Mon Sep 17 00:00:00 2001
From: Joshua Batista <jbatista at microsoft.com>
Date: Tue, 1 Sep 2026 14:32:34 -0700
Subject: [PATCH 2/7] add matrix support, remove enum support
---
.../clang/Basic/DiagnosticSemaKinds.td | 2 ++
clang/include/clang/Basic/HLSLIntrinsics.td | 1 -
clang/lib/Sema/SemaHLSL.cpp | 32 ++++++++++++++++++-
.../builtins/WaveReadLaneFirst.hlsl | 9 ++++++
.../BuiltIns/WaveReadLaneFirst-errors.hlsl | 9 +++++-
.../test/CodeGen/DirectX/WaveReadLaneFirst.ll | 13 ++++++++
6 files changed, 63 insertions(+), 3 deletions(-)
diff --git a/clang/include/clang/Basic/DiagnosticSemaKinds.td b/clang/include/clang/Basic/DiagnosticSemaKinds.td
index 7cb48eaa7d43d4..5909828835066b 100644
--- a/clang/include/clang/Basic/DiagnosticSemaKinds.td
+++ b/clang/include/clang/Basic/DiagnosticSemaKinds.td
@@ -9995,6 +9995,8 @@ def err_typecheck_expect_scalar_or_vector_or_matrix : Error<
"a vector or matrix of such type is required">;
def err_typecheck_expect_any_scalar_or_vector : Error<
"invalid operand of type %0%select{| where a scalar or vector is required}1">;
+def err_typecheck_expect_any_scalar_or_vector_or_matrix : Error<
+ "invalid operand of type %0 where a scalar, vector, or matrix is required">;
def err_typecheck_expect_flt_or_vector : Error<
"invalid operand of type %0 where floating, complex or "
"a vector of such types is required">;
diff --git a/clang/include/clang/Basic/HLSLIntrinsics.td b/clang/include/clang/Basic/HLSLIntrinsics.td
index d4f6fa8c5357d9..cbfee25f75938f 100644
--- a/clang/include/clang/Basic/HLSLIntrinsics.td
+++ b/clang/include/clang/Basic/HLSLIntrinsics.td
@@ -1948,5 +1948,4 @@ def hlsl_wave_read_lane_first :
let VaryingTypes = AllTypesWithBool;
let IsConvergent = 1;
let Availability = SM6_0;
- let VaryingMatDims = [];
}
diff --git a/clang/lib/Sema/SemaHLSL.cpp b/clang/lib/Sema/SemaHLSL.cpp
index 7309dd79f065a0..406bf87eef5b38 100644
--- a/clang/lib/Sema/SemaHLSL.cpp
+++ b/clang/lib/Sema/SemaHLSL.cpp
@@ -3599,6 +3599,35 @@ static bool CheckAnyScalarOrVector(Sema *S, CallExpr *TheCall,
return false;
}
+static bool CheckAnyScalarOrVectorOrMatrix(Sema *S, CallExpr *TheCall,
+ unsigned ArgIndex, bool AllowBool) {
+ assert(TheCall->getNumArgs() > ArgIndex);
+ QualType ArgType = TheCall->getArg(ArgIndex)->getType();
+ if (ArgType->isDependentType())
+ return false;
+
+ QualType ElementType = ArgType;
+ if (const auto *VectorTy = ArgType->getAs<VectorType>())
+ ElementType = VectorTy->getElementType();
+ else if (const auto *MatrixTy = ArgType->getAs<MatrixType>())
+ ElementType = MatrixTy->getElementType();
+
+ if (ElementType->isBooleanType()) {
+ if (AllowBool)
+ return false;
+ } else if ((ElementType->isIntegerType() && !ElementType->isEnumeralType()) ||
+ ElementType->isRealFloatingType()) {
+ unsigned BitWidth = S->Context.getTypeSize(ElementType);
+ if (BitWidth == 16 || BitWidth == 32 || BitWidth == 64)
+ return false;
+ }
+
+ S->Diag(TheCall->getArg(ArgIndex)->getBeginLoc(),
+ diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
+ << ArgType;
+ return true;
+}
+
// Check that the argument is not a bool or vector<bool>
// Returns true on error
static bool CheckNotBoolScalarOrVector(Sema *S, CallExpr *TheCall,
@@ -4833,7 +4862,8 @@ bool SemaHLSL::CheckBuiltinFunctionCall(unsigned BuiltinID, CallExpr *TheCall) {
if (SemaRef.checkArgCount(TheCall, 1))
return true;
- if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
+ if (CheckAnyScalarOrVectorOrMatrix(&SemaRef, TheCall, 0,
+ /*AllowBool=*/true))
return true;
TheCall->setType(TheCall->getArg(0)->getType());
diff --git a/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl b/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl
index 79ddba9d9cc22e..fbae334b327158 100644
--- a/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl
+++ b/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl
@@ -106,4 +106,13 @@ float4 test_floatv4(float4 expr) {
return WaveReadLaneFirst(expr);
}
+// CHECK-LABEL: test_float2x2
+float2x2 test_float2x2(float2x2 expr) {
+ // CHECK-SPIRV: %[[#entry_tok7:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.v4f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok7]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.v4f32([[TY]] %[[#]])
+ // CHECK: ret [[TY]] %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+
// CHECK: attributes #[[#attr]] = {{{.*}} convergent {{.*}}}
diff --git a/clang/test/SemaHLSL/BuiltIns/WaveReadLaneFirst-errors.hlsl b/clang/test/SemaHLSL/BuiltIns/WaveReadLaneFirst-errors.hlsl
index b042362cb79aff..f519200c414b3a 100644
--- a/clang/test/SemaHLSL/BuiltIns/WaveReadLaneFirst-errors.hlsl
+++ b/clang/test/SemaHLSL/BuiltIns/WaveReadLaneFirst-errors.hlsl
@@ -14,5 +14,12 @@ struct S { float f; };
S test_expr_struct_type_check(S p0) {
return __builtin_hlsl_wave_read_lane_first(p0);
- // expected-error at -1 {{invalid operand of type 'S' where a scalar or vector is required}}
+ // expected-error at -1 {{invalid operand of type 'S' where a scalar, vector, or matrix is required}}
+}
+
+enum E { A };
+
+E test_expr_enum_type_check(E p0) {
+ return __builtin_hlsl_wave_read_lane_first(p0);
+ // expected-error at -1 {{invalid operand of type 'E' where a scalar, vector, or matrix is required}}
}
diff --git a/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll b/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll
index 336e4ba90adea9..0a68455aa166e1 100644
--- a/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll
+++ b/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll
@@ -72,6 +72,18 @@ entry:
ret <3 x i32> %ret
}
+define noundef <4 x float> @wave_readlane_first_v4float(
+ <4 x float> noundef %expr) {
+entry:
+; CHECK: call float @dx.op.waveReadLaneFirst.f32(i32 118, float %expr.i0)
+; CHECK: call float @dx.op.waveReadLaneFirst.f32(i32 118, float %expr.i1)
+; CHECK: call float @dx.op.waveReadLaneFirst.f32(i32 118, float %expr.i2)
+; CHECK: call float @dx.op.waveReadLaneFirst.f32(i32 118, float %expr.i3)
+ %ret = call <4 x float> @llvm.dx.wave.readlane.first.v4f32(
+ <4 x float> %expr)
+ ret <4 x float> %ret
+}
+
declare half @llvm.dx.wave.readlane.first.f16(half)
declare float @llvm.dx.wave.readlane.first.f32(float)
declare double @llvm.dx.wave.readlane.first.f64(double)
@@ -81,3 +93,4 @@ declare i32 @llvm.dx.wave.readlane.first.i32(i32)
declare i64 @llvm.dx.wave.readlane.first.i64(i64)
declare <2 x half> @llvm.dx.wave.readlane.first.v2f16(<2 x half>)
declare <3 x i32> @llvm.dx.wave.readlane.first.v3i32(<3 x i32>)
+declare <4 x float> @llvm.dx.wave.readlane.first.v4f32(<4 x float>)
>From 0a049f7c1c5732bddb5a04c56b963bc36cef89ac Mon Sep 17 00:00:00 2001
From: Joshua Batista <jbatista at microsoft.com>
Date: Thu, 10 Sep 2026 14:27:25 -0700
Subject: [PATCH 3/7] address Deric
---
.../clang/Basic/DiagnosticSemaKinds.td | 5 ++--
clang/lib/Sema/SemaHLSL.cpp | 26 +++++++++----------
.../BuiltIns/WaveReadLaneFirst-errors.hlsl | 1 -
3 files changed, 14 insertions(+), 18 deletions(-)
diff --git a/clang/include/clang/Basic/DiagnosticSemaKinds.td b/clang/include/clang/Basic/DiagnosticSemaKinds.td
index 5909828835066b..be3ff0819c7436 100644
--- a/clang/include/clang/Basic/DiagnosticSemaKinds.td
+++ b/clang/include/clang/Basic/DiagnosticSemaKinds.td
@@ -9993,10 +9993,9 @@ def err_typecheck_expect_scalar_or_vector : Error<
def err_typecheck_expect_scalar_or_vector_or_matrix : Error<
"invalid operand of type %0 where %1 or "
"a vector or matrix of such type is required">;
-def err_typecheck_expect_any_scalar_or_vector : Error<
- "invalid operand of type %0%select{| where a scalar or vector is required}1">;
def err_typecheck_expect_any_scalar_or_vector_or_matrix : Error<
- "invalid operand of type %0 where a scalar, vector, or matrix is required">;
+ "invalid operand of type %0%select{| where a scalar or vector is required|"
+ " where a scalar, vector, or matrix is required}1">;
def err_typecheck_expect_flt_or_vector : Error<
"invalid operand of type %0 where floating, complex or "
"a vector of such types is required">;
diff --git a/clang/lib/Sema/SemaHLSL.cpp b/clang/lib/Sema/SemaHLSL.cpp
index 406bf87eef5b38..31b4700f919915 100644
--- a/clang/lib/Sema/SemaHLSL.cpp
+++ b/clang/lib/Sema/SemaHLSL.cpp
@@ -3592,7 +3592,7 @@ static bool CheckAnyScalarOrVector(Sema *S, CallExpr *TheCall,
if (!(ArgType->isScalarType() ||
(VTy && VTy->getElementType()->isScalarType()))) {
S->Diag(TheCall->getArg(0)->getBeginLoc(),
- diag::err_typecheck_expect_any_scalar_or_vector)
+ diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
<< ArgType << 1;
return true;
}
@@ -3600,7 +3600,7 @@ static bool CheckAnyScalarOrVector(Sema *S, CallExpr *TheCall,
}
static bool CheckAnyScalarOrVectorOrMatrix(Sema *S, CallExpr *TheCall,
- unsigned ArgIndex, bool AllowBool) {
+ unsigned ArgIndex) {
assert(TheCall->getNumArgs() > ArgIndex);
QualType ArgType = TheCall->getArg(ArgIndex)->getType();
if (ArgType->isDependentType())
@@ -3609,14 +3609,13 @@ static bool CheckAnyScalarOrVectorOrMatrix(Sema *S, CallExpr *TheCall,
QualType ElementType = ArgType;
if (const auto *VectorTy = ArgType->getAs<VectorType>())
ElementType = VectorTy->getElementType();
- else if (const auto *MatrixTy = ArgType->getAs<MatrixType>())
+ else if (const auto *MatrixTy = ArgType->getAs<ConstantMatrixType>())
ElementType = MatrixTy->getElementType();
- if (ElementType->isBooleanType()) {
- if (AllowBool)
- return false;
- } else if ((ElementType->isIntegerType() && !ElementType->isEnumeralType()) ||
- ElementType->isRealFloatingType()) {
+ if (ElementType->isBooleanType())
+ return false;
+
+ if (ElementType->isIntegerType() || ElementType->isRealFloatingType()) {
unsigned BitWidth = S->Context.getTypeSize(ElementType);
if (BitWidth == 16 || BitWidth == 32 || BitWidth == 64)
return false;
@@ -3624,7 +3623,7 @@ static bool CheckAnyScalarOrVectorOrMatrix(Sema *S, CallExpr *TheCall,
S->Diag(TheCall->getArg(ArgIndex)->getBeginLoc(),
diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
- << ArgType;
+ << ArgType << 2;
return true;
}
@@ -3641,7 +3640,7 @@ static bool CheckNotBoolScalarOrVector(Sema *S, CallExpr *TheCall,
(VTy &&
S->Context.hasSameUnqualifiedType(VTy->getElementType(), BoolType))) {
S->Diag(TheCall->getArg(0)->getBeginLoc(),
- diag::err_typecheck_expect_any_scalar_or_vector)
+ diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
<< ArgType << 0;
return true;
}
@@ -4821,14 +4820,14 @@ bool SemaHLSL::CheckBuiltinFunctionCall(unsigned BuiltinID, CallExpr *TheCall) {
if (!(ArgType->isScalarType())) {
SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
- diag::err_typecheck_expect_any_scalar_or_vector)
+ diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
<< ArgType << 0;
return true;
}
if (!(ArgType->isBooleanType())) {
SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
- diag::err_typecheck_expect_any_scalar_or_vector)
+ diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
<< ArgType << 0;
return true;
}
@@ -4862,8 +4861,7 @@ bool SemaHLSL::CheckBuiltinFunctionCall(unsigned BuiltinID, CallExpr *TheCall) {
if (SemaRef.checkArgCount(TheCall, 1))
return true;
- if (CheckAnyScalarOrVectorOrMatrix(&SemaRef, TheCall, 0,
- /*AllowBool=*/true))
+ if (CheckAnyScalarOrVectorOrMatrix(&SemaRef, TheCall, 0))
return true;
TheCall->setType(TheCall->getArg(0)->getType());
diff --git a/clang/test/SemaHLSL/BuiltIns/WaveReadLaneFirst-errors.hlsl b/clang/test/SemaHLSL/BuiltIns/WaveReadLaneFirst-errors.hlsl
index f519200c414b3a..19e5fdadc276df 100644
--- a/clang/test/SemaHLSL/BuiltIns/WaveReadLaneFirst-errors.hlsl
+++ b/clang/test/SemaHLSL/BuiltIns/WaveReadLaneFirst-errors.hlsl
@@ -21,5 +21,4 @@ enum E { A };
E test_expr_enum_type_check(E p0) {
return __builtin_hlsl_wave_read_lane_first(p0);
- // expected-error at -1 {{invalid operand of type 'E' where a scalar, vector, or matrix is required}}
}
>From d9338a7c14ffdec7eb386c7e39a85e99c0ca6f78 Mon Sep 17 00:00:00 2001
From: Joshua Batista <jbatista at microsoft.com>
Date: Fri, 18 Sep 2026 15:38:26 -0700
Subject: [PATCH 4/7] attempt to address Farzon
---
clang/include/clang/Basic/HLSLIntrinsics.td | 5 +-
.../builtins/WaveReadLaneFirst.hlsl | 68 ++++++++++++-------
.../CodeGen/GlobalISel/LegalizerHelper.cpp | 20 +++++-
llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp | 13 ++++
llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp | 11 +--
.../test/CodeGen/DirectX/WaveReadLaneFirst.ll | 33 +++++++++
.../hlsl-intrinsics/WaveReadLaneFirst.ll | 45 +++++++++++-
7 files changed, 160 insertions(+), 35 deletions(-)
diff --git a/clang/include/clang/Basic/HLSLIntrinsics.td b/clang/include/clang/Basic/HLSLIntrinsics.td
index cbfee25f75938f..1b079bc90dfa6c 100644
--- a/clang/include/clang/Basic/HLSLIntrinsics.td
+++ b/clang/include/clang/Basic/HLSLIntrinsics.td
@@ -1946,6 +1946,7 @@ def hlsl_wave_read_lane_first :
\param Val The value to read.
}];
let VaryingTypes = AllTypesWithBool;
- let IsConvergent = 1;
- let Availability = SM6_0;
+let VaryingLongVector = 1;
+let IsConvergent = 1;
+let Availability = SM6_0;
}
diff --git a/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl b/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl
index fbae334b327158..644440a6bcfc46 100644
--- a/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl
+++ b/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl
@@ -7,9 +7,9 @@
// CHECK-LABEL: test_int
int test_int(int expr) {
- // CHECK-SPIRV: %[[#entry_tok0:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK: %[[#entry_tok0:]] = call token @llvm.experimental.convergence.entry()
// CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok0]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i32([[TY]] %[[#]])
+ // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok0]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -19,18 +19,18 @@ int test_int(int expr) {
// CHECK-LABEL: test_uint
uint test_uint(uint expr) {
- // CHECK-SPIRV: %[[#entry_tok0:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK: %[[#entry_tok0:]] = call token @llvm.experimental.convergence.entry()
// CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok0]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i32([[TY]] %[[#]])
+ // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok0]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
// CHECK-LABEL: test_int64_t
int64_t test_int64_t(int64_t expr) {
- // CHECK-SPIRV: %[[#entry_tok1:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK: %[[#entry_tok1:]] = call token @llvm.experimental.convergence.entry()
// CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok1]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i64([[TY]] %[[#]])
+ // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok1]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -40,9 +40,9 @@ int64_t test_int64_t(int64_t expr) {
// CHECK-LABEL: test_uint64_t
uint64_t test_uint64_t(uint64_t expr) {
- // CHECK-SPIRV: %[[#entry_tok1:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK: %[[#entry_tok1:]] = call token @llvm.experimental.convergence.entry()
// CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok1]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i64([[TY]] %[[#]])
+ // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok1]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -50,9 +50,9 @@ uint64_t test_uint64_t(uint64_t expr) {
#ifdef __HLSL_ENABLE_16_BIT
// CHECK-LABEL: test_int16
int16_t test_int16(int16_t expr) {
- // CHECK-SPIRV: %[[#entry_tok2:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK: %[[#entry_tok2:]] = call token @llvm.experimental.convergence.entry()
// CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok2]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i16([[TY]] %[[#]])
+ // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok2]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -62,9 +62,9 @@ int16_t test_int16(int16_t expr) {
// CHECK-LABEL: test_uint16
uint16_t test_uint16(uint16_t expr) {
- // CHECK-SPIRV: %[[#entry_tok2:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK: %[[#entry_tok2:]] = call token @llvm.experimental.convergence.entry()
// CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok2]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i16([[TY]] %[[#]])
+ // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok2]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -72,45 +72,63 @@ uint16_t test_uint16(uint16_t expr) {
// CHECK-LABEL: test_bool
bool test_bool(bool expr) {
- // CHECK-SPIRV: %[[#entry_tok3:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK: %[[#entry_tok3:]] = call token @llvm.experimental.convergence.entry()
// CHECK-SPIRV: %[[RET:.*]] = call i1 @llvm.spv.wave.readlane.first.i1(i1 %{{[a-zA-Z0-9]+}}) [ "convergencectrl"(token %[[#entry_tok3]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call i1 @llvm.dx.wave.readlane.first.i1(i1 %{{[a-zA-Z0-9]+}})
+ // CHECK-DXIL: %[[RET:.*]] = call i1 @llvm.dx.wave.readlane.first.i1(i1 %{{[a-zA-Z0-9]+}}) [ "convergencectrl"(token %[[#entry_tok3]]) ]
// CHECK: ret i1 %[[RET]]
return WaveReadLaneFirst(expr);
}
// CHECK-LABEL: test_half
half test_half(half expr) {
- // CHECK-SPIRV: %[[#entry_tok4:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK: %[[#entry_tok4:]] = call token @llvm.experimental.convergence.entry()
// CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.f16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok4]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.f16([[TY]] %[[#]])
+ // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.f16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok4]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
// CHECK-LABEL: test_double
double test_double(double expr) {
- // CHECK-SPIRV: %[[#entry_tok5:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK: %[[#entry_tok5:]] = call token @llvm.experimental.convergence.entry()
// CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.f64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok5]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.f64([[TY]] %[[#]])
+ // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.f64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok5]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
// CHECK-LABEL: test_floatv4
float4 test_floatv4(float4 expr) {
- // CHECK-SPIRV: %[[#entry_tok6:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK: %[[#entry_tok6:]] = call token @llvm.experimental.convergence.entry()
// CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.v4f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok6]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.v4f32([[TY]] %[[#]])
+ // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.v4f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok6]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
-// CHECK-LABEL: test_float2x2
-float2x2 test_float2x2(float2x2 expr) {
- // CHECK-SPIRV: %[[#entry_tok7:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.v4f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok7]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.v4f32([[TY]] %[[#]])
+// CHECK-LABEL: test_floatv5
+vector<float, 5> test_floatv5(vector<float, 5> expr) {
+ // CHECK: %[[#entry_tok7:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.v5f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok7]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.v5f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok7]]) ]
+ // CHECK: ret [[TY]] %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+
+// CHECK-LABEL: test_float2x3
+float2x3 test_float2x3(float2x3 expr) {
+ // CHECK: %[[#entry_tok8:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.v6f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok8]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.v6f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok8]]) ]
+ // CHECK: ret [[TY]] %[[RET]]
+ return WaveReadLaneFirst(expr);
+}
+
+// CHECK-LABEL: test_float3x4
+float3x4 test_float3x4(float3x4 expr) {
+ // CHECK: %[[#entry_tok9:]] = call token @llvm.experimental.convergence.entry()
+ // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.v12f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok9]]) ]
+ // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.v12f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok9]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
diff --git a/llvm/lib/CodeGen/GlobalISel/LegalizerHelper.cpp b/llvm/lib/CodeGen/GlobalISel/LegalizerHelper.cpp
index 4ab02b0b12bff5..9f7cd14d4fcdf9 100644
--- a/llvm/lib/CodeGen/GlobalISel/LegalizerHelper.cpp
+++ b/llvm/lib/CodeGen/GlobalISel/LegalizerHelper.cpp
@@ -5230,8 +5230,11 @@ static bool hasSameNumEltsOnAllVectorOperands(
if (!VecTy.isVector())
return false;
unsigned NumElts = VecTy.getNumElements();
+ unsigned IntrinsicIDOp = isa<GIntrinsic>(MI) ? MI.getNumExplicitDefs() : ~0U;
for (unsigned OpIdx = 1; OpIdx < MI.getNumOperands(); ++OpIdx) {
+ if (OpIdx == IntrinsicIDOp)
+ continue;
MachineOperand &Op = MI.getOperand(OpIdx);
if (!Op.isReg()) {
if (!is_contained(NonVecOpIndices, OpIdx))
@@ -5314,8 +5317,10 @@ LegalizerHelper::fewerElementsVectorMultiEltType(
"Non-compatible opcode or not specified non-vector operands");
unsigned OrigNumElts = MRI.getType(MI.getReg(0)).getNumElements();
- unsigned NumInputs = MI.getNumOperands() - MI.getNumDefs();
unsigned NumDefs = MI.getNumDefs();
+ auto *GI = dyn_cast<GIntrinsic>(&MI);
+ unsigned FirstUse = NumDefs + (GI ? 1 : 0);
+ unsigned NumInputs = MI.getNumOperands() - FirstUse;
// Create DstOps (sub-vectors with NumElts elts + Leftover) for each output.
// Build instructions with DstOps to use instruction found by CSE directly.
@@ -5332,7 +5337,7 @@ LegalizerHelper::fewerElementsVectorMultiEltType(
// examples: compare predicate in icmp and fcmp (op 1), vector select with i1
// scalar condition (op 1), immediate in sext_inreg (op 2).
SmallVector<SmallVector<SrcOp, 8>, 3> InputOpsPieces(NumInputs);
- for (unsigned UseIdx = NumDefs, UseNo = 0; UseIdx < MI.getNumOperands();
+ for (unsigned UseIdx = FirstUse, UseNo = 0; UseIdx < MI.getNumOperands();
++UseIdx, ++UseNo) {
if (is_contained(NonVecOpIndices, UseIdx)) {
broadcastSrcOp(InputOpsPieces[UseNo], OutputOpsPieces[0].size(),
@@ -5358,7 +5363,16 @@ LegalizerHelper::fewerElementsVectorMultiEltType(
for (unsigned InputNo = 0; InputNo < NumInputs; ++InputNo)
Uses.push_back(InputOpsPieces[InputNo][i]);
- auto I = MIRBuilder.buildInstr(MI.getOpcode(), Defs, Uses, MI.getFlags());
+ MachineInstrBuilder I;
+ if (GI) {
+ I = MIRBuilder.buildIntrinsic(GI->getIntrinsicID(), Defs,
+ GI->hasSideEffects(), GI->isConvergent());
+ I.setMIFlags(MI.getFlags());
+ for (SrcOp &Use : Uses)
+ Use.addSrcToMIB(I);
+ } else {
+ I = MIRBuilder.buildInstr(MI.getOpcode(), Defs, Uses, MI.getFlags());
+ }
for (unsigned DstNo = 0; DstNo < NumDefs; ++DstNo)
OutputRegs[DstNo].push_back(I.getReg(DstNo));
}
diff --git a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
index dcd89eb3a69d54..484e18fb9cada4 100644
--- a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
@@ -1040,6 +1040,19 @@ bool SPIRVLegalizerInfo::legalizeIntrinsic(LegalizerHelper &Helper,
MachineInstr &MI) const {
LLVM_DEBUG(dbgs() << "legalizeIntrinsic: " << MI);
auto IntrinsicID = cast<GIntrinsic>(MI).getIntrinsicID();
+
+ if (IntrinsicID == Intrinsic::spv_wave_readlane_first) {
+ MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
+ LLT DstTy = MRI.getType(MI.getOperand(0).getReg());
+ if (needsVectorLegalization(DstTy, *ST)) {
+ unsigned MaxVectorSize = ST->isShader() ? 4 : 16;
+ unsigned NumElts = llvm::bit_floor(
+ std::min<unsigned>(DstTy.getNumElements(), MaxVectorSize));
+ return Helper.fewerElementsVectorMultiEltType(
+ cast<GIntrinsic>(MI), NumElts) == LegalizerHelper::Legalized;
+ }
+ }
+
switch (IntrinsicID) {
case Intrinsic::spv_bitcast:
return legalizeSpvBitcast(Helper, MI, GR);
diff --git a/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp b/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp
index 63f6bf5f4483d9..e990963210f68a 100644
--- a/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp
@@ -187,12 +187,12 @@ static SPIRVTypeInst deduceTypeFromUses(Register Reg, MachineFunction &MF,
ResType = deduceTypeFromPointerOperand(&Use, Reg, GR, MIB);
break;
case TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS:
+ case TargetOpcode::G_INTRINSIC_CONVERGENT:
case TargetOpcode::G_INTRINSIC: {
auto IntrinsicID = cast<GIntrinsic>(Use).getIntrinsicID();
- if (IntrinsicID == Intrinsic::spv_insertelt) {
- if (Reg == Use.getOperand(2).getReg())
- ResType = deduceTypeFromResultRegister(&Use, Reg, GR, MIB);
- } else if (IntrinsicID == Intrinsic::spv_extractelt) {
+ if (IntrinsicID == Intrinsic::spv_wave_readlane_first ||
+ IntrinsicID == Intrinsic::spv_insertelt ||
+ IntrinsicID == Intrinsic::spv_extractelt) {
if (Reg == Use.getOperand(2).getReg())
ResType = deduceTypeFromResultRegister(&Use, Reg, GR, MIB);
}
@@ -296,10 +296,13 @@ static SPIRVTypeInst deduceResultTypeFromOperands(MachineInstr *I,
case TargetOpcode::G_SHUFFLE_VECTOR:
return deduceTypeFromOperandRange(I, MIB, GR, 1, 3);
case TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS:
+ case TargetOpcode::G_INTRINSIC_CONVERGENT:
case TargetOpcode::G_INTRINSIC: {
auto IntrinsicID = cast<GIntrinsic>(I)->getIntrinsicID();
if (IntrinsicID == Intrinsic::spv_gep)
return deduceGEPType(I, GR, MIB);
+ if (IntrinsicID == Intrinsic::spv_wave_readlane_first)
+ return deduceTypeFromSingleOperand(I, MIB, GR, 2);
break;
}
case TargetOpcode::G_LOAD: {
diff --git a/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll b/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll
index 0a68455aa166e1..c890eef2533472 100644
--- a/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll
+++ b/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll
@@ -84,6 +84,36 @@ entry:
ret <4 x float> %ret
}
+define noundef <5 x float> @wave_readlane_first_v5float(
+ <5 x float> noundef %expr) {
+entry:
+; CHECK-LABEL: define noundef <5 x float> @wave_readlane_first_v5float(
+; CHECK-COUNT-5: call float @dx.op.waveReadLaneFirst.f32(i32 118,
+ %ret = call <5 x float> @llvm.dx.wave.readlane.first.v5f32(
+ <5 x float> %expr)
+ ret <5 x float> %ret
+}
+
+define noundef <6 x float> @wave_readlane_first_float2x3(
+ <6 x float> noundef %expr) {
+entry:
+; CHECK-LABEL: define noundef <6 x float> @wave_readlane_first_float2x3(
+; CHECK-COUNT-6: call float @dx.op.waveReadLaneFirst.f32(i32 118,
+ %ret = call <6 x float> @llvm.dx.wave.readlane.first.v6f32(
+ <6 x float> %expr)
+ ret <6 x float> %ret
+}
+
+define noundef <12 x float> @wave_readlane_first_float3x4(
+ <12 x float> noundef %expr) {
+entry:
+; CHECK-LABEL: define noundef <12 x float> @wave_readlane_first_float3x4(
+; CHECK-COUNT-12: call float @dx.op.waveReadLaneFirst.f32(i32 118,
+ %ret = call <12 x float> @llvm.dx.wave.readlane.first.v12f32(
+ <12 x float> %expr)
+ ret <12 x float> %ret
+}
+
declare half @llvm.dx.wave.readlane.first.f16(half)
declare float @llvm.dx.wave.readlane.first.f32(float)
declare double @llvm.dx.wave.readlane.first.f64(double)
@@ -94,3 +124,6 @@ declare i64 @llvm.dx.wave.readlane.first.i64(i64)
declare <2 x half> @llvm.dx.wave.readlane.first.v2f16(<2 x half>)
declare <3 x i32> @llvm.dx.wave.readlane.first.v3i32(<3 x i32>)
declare <4 x float> @llvm.dx.wave.readlane.first.v4f32(<4 x float>)
+declare <5 x float> @llvm.dx.wave.readlane.first.v5f32(<5 x float>)
+declare <6 x float> @llvm.dx.wave.readlane.first.v6f32(<6 x float>)
+declare <12 x float> @llvm.dx.wave.readlane.first.v12f32(<12 x float>)
diff --git a/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
index 67651fee276973..db23472353e8fb 100644
--- a/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
+++ b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
@@ -1,17 +1,22 @@
; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv1.5-vulkan-unknown %s -o - | FileCheck %s
; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv1.5-vulkan-unknown %s -o - -filetype=obj | spirv-val %}
-; Test WaveReadLaneFirst lowering for scalar and vector types.
+; Test WaveReadLaneFirst lowering for scalar, vector, and matrix types.
; CHECK: Capability Shader
; CHECK: Capability GroupNonUniformBallot
; CHECK-DAG: %[[#uint:]] = OpTypeInt 32 0
; CHECK-DAG: %[[#f32:]] = OpTypeFloat 32
+; CHECK-DAG: %[[#v2_float:]] = OpTypeVector %[[#f32]] 2
; CHECK-DAG: %[[#v4_float:]] = OpTypeVector %[[#f32]] 4
; CHECK-DAG: %[[#bool:]] = OpTypeBool
; CHECK-DAG: %[[#scope:]] = OpConstant %[[#uint]] 3
+ at wide_f32_5 = internal addrspace(10) global [5 x float] zeroinitializer
+ at wide_f32_6 = internal addrspace(10) global [6 x float] zeroinitializer
+ at wide_f32_12 = internal addrspace(10) global [12 x float] zeroinitializer
+
; CHECK-LABEL: Begin function test_float
; CHECK: %[[#fexpr:]] = OpFunctionParameter %[[#f32]]
define float @test_float(float %fexpr) {
@@ -49,7 +54,45 @@ entry:
ret <4 x float> %0
}
+; CHECK-LABEL: Begin function test_floatv5
+define void @test_floatv5() {
+entry:
+ %expr = load <5 x float>, ptr addrspace(10) @wide_f32_5
+; CHECK: OpGroupNonUniformBroadcastFirst %[[#v4_float]] %[[#scope]]
+; CHECK: OpGroupNonUniformBroadcastFirst %[[#f32]] %[[#scope]]
+ %result = call <5 x float> @llvm.spv.wave.readlane.first.v5f32(
+ <5 x float> %expr)
+ store <5 x float> %result, ptr addrspace(10) @wide_f32_5
+ ret void
+}
+
+; CHECK-LABEL: Begin function test_float2x3
+define void @test_float2x3() {
+entry:
+ %expr = load <6 x float>, ptr addrspace(10) @wide_f32_6
+; CHECK: OpGroupNonUniformBroadcastFirst %[[#v4_float]] %[[#scope]]
+; CHECK: OpGroupNonUniformBroadcastFirst %[[#v2_float]] %[[#scope]]
+ %result = call <6 x float> @llvm.spv.wave.readlane.first.v6f32(
+ <6 x float> %expr)
+ store <6 x float> %result, ptr addrspace(10) @wide_f32_6
+ ret void
+}
+
+; CHECK-LABEL: Begin function test_float3x4
+define void @test_float3x4() {
+entry:
+ %expr = load <12 x float>, ptr addrspace(10) @wide_f32_12
+; CHECK-COUNT-3: OpGroupNonUniformBroadcastFirst %[[#v4_float]] %[[#scope]]
+ %result = call <12 x float> @llvm.spv.wave.readlane.first.v12f32(
+ <12 x float> %expr)
+ store <12 x float> %result, ptr addrspace(10) @wide_f32_12
+ ret void
+}
+
declare float @llvm.spv.wave.readlane.first.f32(float)
declare i32 @llvm.spv.wave.readlane.first.i32(i32)
declare i1 @llvm.spv.wave.readlane.first.i1(i1)
declare <4 x float> @llvm.spv.wave.readlane.first.v4f32(<4 x float>)
+declare <5 x float> @llvm.spv.wave.readlane.first.v5f32(<5 x float>)
+declare <6 x float> @llvm.spv.wave.readlane.first.v6f32(<6 x float>)
+declare <12 x float> @llvm.spv.wave.readlane.first.v12f32(<12 x float>)
>From 86a23526d2d8fa542bd7968821b7bce9ab4891d3 Mon Sep 17 00:00:00 2001
From: Joshua Batista <jbatista at microsoft.com>
Date: Fri, 18 Sep 2026 15:58:28 -0700
Subject: [PATCH 5/7] touch ups
---
.../builtins/WaveReadLaneFirst.hlsl | 52 +++++++------------
.../CodeGen/GlobalISel/LegalizerHelper.cpp | 20 ++-----
llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp | 13 -----
llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp | 11 ++--
.../hlsl-intrinsics/WaveReadLaneFirst.ll | 21 +++++---
5 files changed, 38 insertions(+), 79 deletions(-)
diff --git a/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl b/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl
index 644440a6bcfc46..fb3bd5bbb042ef 100644
--- a/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl
+++ b/clang/test/CodeGenHLSL/builtins/WaveReadLaneFirst.hlsl
@@ -1,27 +1,24 @@
// RUN: %clang_cc1 -std=hlsl2021 -finclude-default-header -fnative-half-type -fnative-int16-type -triple \
// RUN: dxil-pc-shadermodel6.3-library %s -emit-llvm -disable-llvm-passes -o - | \
-// RUN: FileCheck %s --check-prefixes=CHECK,CHECK-DXIL
+// RUN: FileCheck %s -DTARGET=dx
// RUN: %clang_cc1 -std=hlsl2021 -finclude-default-header -fnative-half-type -fnative-int16-type -triple \
// RUN: spirv-pc-vulkan-library %s -emit-llvm -disable-llvm-passes -o - | \
-// RUN: FileCheck %s --check-prefixes=CHECK,CHECK-SPIRV
+// RUN: FileCheck %s -DTARGET=spv
// CHECK-LABEL: test_int
int test_int(int expr) {
// CHECK: %[[#entry_tok0:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok0]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok0]]) ]
+ // CHECK: %[[RET:.*]] = call [[TY:.*]] @llvm.[[TARGET]].wave.readlane.first.i32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok0]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
-// CHECK-DXIL: declare [[TY]] @llvm.dx.wave.readlane.first.i32([[TY]]) #[[#attr:]]
-// CHECK-SPIRV: declare [[TY]] @llvm.spv.wave.readlane.first.i32([[TY]]) #[[#attr:]]
+// CHECK: declare [[TY]] @llvm.[[TARGET]].wave.readlane.first.i32([[TY]]) #[[#attr:]]
// CHECK-LABEL: test_uint
uint test_uint(uint expr) {
// CHECK: %[[#entry_tok0:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok0]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok0]]) ]
+ // CHECK: %[[RET:.*]] = call [[TY:.*]] @llvm.[[TARGET]].wave.readlane.first.i32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok0]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -29,20 +26,17 @@ uint test_uint(uint expr) {
// CHECK-LABEL: test_int64_t
int64_t test_int64_t(int64_t expr) {
// CHECK: %[[#entry_tok1:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok1]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok1]]) ]
+ // CHECK: %[[RET:.*]] = call [[TY:.*]] @llvm.[[TARGET]].wave.readlane.first.i64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok1]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
-// CHECK-DXIL: declare [[TY]] @llvm.dx.wave.readlane.first.i64([[TY]]) #[[#attr:]]
-// CHECK-SPIRV: declare [[TY]] @llvm.spv.wave.readlane.first.i64([[TY]]) #[[#attr:]]
+// CHECK: declare [[TY]] @llvm.[[TARGET]].wave.readlane.first.i64([[TY]]) #[[#attr:]]
// CHECK-LABEL: test_uint64_t
uint64_t test_uint64_t(uint64_t expr) {
// CHECK: %[[#entry_tok1:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok1]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok1]]) ]
+ // CHECK: %[[RET:.*]] = call [[TY:.*]] @llvm.[[TARGET]].wave.readlane.first.i64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok1]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -51,20 +45,17 @@ uint64_t test_uint64_t(uint64_t expr) {
// CHECK-LABEL: test_int16
int16_t test_int16(int16_t expr) {
// CHECK: %[[#entry_tok2:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok2]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok2]]) ]
+ // CHECK: %[[RET:.*]] = call [[TY:.*]] @llvm.[[TARGET]].wave.readlane.first.i16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok2]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
-// CHECK-DXIL: declare [[TY]] @llvm.dx.wave.readlane.first.i16([[TY]]) #[[#attr:]]
-// CHECK-SPIRV: declare [[TY]] @llvm.spv.wave.readlane.first.i16([[TY]]) #[[#attr:]]
+// CHECK: declare [[TY]] @llvm.[[TARGET]].wave.readlane.first.i16([[TY]]) #[[#attr:]]
// CHECK-LABEL: test_uint16
uint16_t test_uint16(uint16_t expr) {
// CHECK: %[[#entry_tok2:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call [[TY:.*]] @llvm.spv.wave.readlane.first.i16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok2]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call [[TY:.*]] @llvm.dx.wave.readlane.first.i16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok2]]) ]
+ // CHECK: %[[RET:.*]] = call [[TY:.*]] @llvm.[[TARGET]].wave.readlane.first.i16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok2]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -73,8 +64,7 @@ uint16_t test_uint16(uint16_t expr) {
// CHECK-LABEL: test_bool
bool test_bool(bool expr) {
// CHECK: %[[#entry_tok3:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call i1 @llvm.spv.wave.readlane.first.i1(i1 %{{[a-zA-Z0-9]+}}) [ "convergencectrl"(token %[[#entry_tok3]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call i1 @llvm.dx.wave.readlane.first.i1(i1 %{{[a-zA-Z0-9]+}}) [ "convergencectrl"(token %[[#entry_tok3]]) ]
+ // CHECK: %[[RET:.*]] = call i1 @llvm.[[TARGET]].wave.readlane.first.i1(i1 %{{[a-zA-Z0-9]+}}) [ "convergencectrl"(token %[[#entry_tok3]]) ]
// CHECK: ret i1 %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -82,8 +72,7 @@ bool test_bool(bool expr) {
// CHECK-LABEL: test_half
half test_half(half expr) {
// CHECK: %[[#entry_tok4:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.f16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok4]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.f16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok4]]) ]
+ // CHECK: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.[[TARGET]].wave.readlane.first.f16([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok4]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -91,8 +80,7 @@ half test_half(half expr) {
// CHECK-LABEL: test_double
double test_double(double expr) {
// CHECK: %[[#entry_tok5:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.f64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok5]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.f64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok5]]) ]
+ // CHECK: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.[[TARGET]].wave.readlane.first.f64([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok5]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -100,8 +88,7 @@ double test_double(double expr) {
// CHECK-LABEL: test_floatv4
float4 test_floatv4(float4 expr) {
// CHECK: %[[#entry_tok6:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.v4f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok6]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.v4f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok6]]) ]
+ // CHECK: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.[[TARGET]].wave.readlane.first.v4f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok6]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -109,8 +96,7 @@ float4 test_floatv4(float4 expr) {
// CHECK-LABEL: test_floatv5
vector<float, 5> test_floatv5(vector<float, 5> expr) {
// CHECK: %[[#entry_tok7:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.v5f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok7]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.v5f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok7]]) ]
+ // CHECK: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.[[TARGET]].wave.readlane.first.v5f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok7]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -118,8 +104,7 @@ vector<float, 5> test_floatv5(vector<float, 5> expr) {
// CHECK-LABEL: test_float2x3
float2x3 test_float2x3(float2x3 expr) {
// CHECK: %[[#entry_tok8:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.v6f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok8]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.v6f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok8]]) ]
+ // CHECK: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.[[TARGET]].wave.readlane.first.v6f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok8]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
@@ -127,8 +112,7 @@ float2x3 test_float2x3(float2x3 expr) {
// CHECK-LABEL: test_float3x4
float3x4 test_float3x4(float3x4 expr) {
// CHECK: %[[#entry_tok9:]] = call token @llvm.experimental.convergence.entry()
- // CHECK-SPIRV: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.spv.wave.readlane.first.v12f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok9]]) ]
- // CHECK-DXIL: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.dx.wave.readlane.first.v12f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok9]]) ]
+ // CHECK: %[[RET:.*]] = call reassoc nnan ninf nsz arcp afn [[TY:.*]] @llvm.[[TARGET]].wave.readlane.first.v12f32([[TY]] %[[#]]) [ "convergencectrl"(token %[[#entry_tok9]]) ]
// CHECK: ret [[TY]] %[[RET]]
return WaveReadLaneFirst(expr);
}
diff --git a/llvm/lib/CodeGen/GlobalISel/LegalizerHelper.cpp b/llvm/lib/CodeGen/GlobalISel/LegalizerHelper.cpp
index 9f7cd14d4fcdf9..4ab02b0b12bff5 100644
--- a/llvm/lib/CodeGen/GlobalISel/LegalizerHelper.cpp
+++ b/llvm/lib/CodeGen/GlobalISel/LegalizerHelper.cpp
@@ -5230,11 +5230,8 @@ static bool hasSameNumEltsOnAllVectorOperands(
if (!VecTy.isVector())
return false;
unsigned NumElts = VecTy.getNumElements();
- unsigned IntrinsicIDOp = isa<GIntrinsic>(MI) ? MI.getNumExplicitDefs() : ~0U;
for (unsigned OpIdx = 1; OpIdx < MI.getNumOperands(); ++OpIdx) {
- if (OpIdx == IntrinsicIDOp)
- continue;
MachineOperand &Op = MI.getOperand(OpIdx);
if (!Op.isReg()) {
if (!is_contained(NonVecOpIndices, OpIdx))
@@ -5317,10 +5314,8 @@ LegalizerHelper::fewerElementsVectorMultiEltType(
"Non-compatible opcode or not specified non-vector operands");
unsigned OrigNumElts = MRI.getType(MI.getReg(0)).getNumElements();
+ unsigned NumInputs = MI.getNumOperands() - MI.getNumDefs();
unsigned NumDefs = MI.getNumDefs();
- auto *GI = dyn_cast<GIntrinsic>(&MI);
- unsigned FirstUse = NumDefs + (GI ? 1 : 0);
- unsigned NumInputs = MI.getNumOperands() - FirstUse;
// Create DstOps (sub-vectors with NumElts elts + Leftover) for each output.
// Build instructions with DstOps to use instruction found by CSE directly.
@@ -5337,7 +5332,7 @@ LegalizerHelper::fewerElementsVectorMultiEltType(
// examples: compare predicate in icmp and fcmp (op 1), vector select with i1
// scalar condition (op 1), immediate in sext_inreg (op 2).
SmallVector<SmallVector<SrcOp, 8>, 3> InputOpsPieces(NumInputs);
- for (unsigned UseIdx = FirstUse, UseNo = 0; UseIdx < MI.getNumOperands();
+ for (unsigned UseIdx = NumDefs, UseNo = 0; UseIdx < MI.getNumOperands();
++UseIdx, ++UseNo) {
if (is_contained(NonVecOpIndices, UseIdx)) {
broadcastSrcOp(InputOpsPieces[UseNo], OutputOpsPieces[0].size(),
@@ -5363,16 +5358,7 @@ LegalizerHelper::fewerElementsVectorMultiEltType(
for (unsigned InputNo = 0; InputNo < NumInputs; ++InputNo)
Uses.push_back(InputOpsPieces[InputNo][i]);
- MachineInstrBuilder I;
- if (GI) {
- I = MIRBuilder.buildIntrinsic(GI->getIntrinsicID(), Defs,
- GI->hasSideEffects(), GI->isConvergent());
- I.setMIFlags(MI.getFlags());
- for (SrcOp &Use : Uses)
- Use.addSrcToMIB(I);
- } else {
- I = MIRBuilder.buildInstr(MI.getOpcode(), Defs, Uses, MI.getFlags());
- }
+ auto I = MIRBuilder.buildInstr(MI.getOpcode(), Defs, Uses, MI.getFlags());
for (unsigned DstNo = 0; DstNo < NumDefs; ++DstNo)
OutputRegs[DstNo].push_back(I.getReg(DstNo));
}
diff --git a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
index 484e18fb9cada4..dcd89eb3a69d54 100644
--- a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
@@ -1040,19 +1040,6 @@ bool SPIRVLegalizerInfo::legalizeIntrinsic(LegalizerHelper &Helper,
MachineInstr &MI) const {
LLVM_DEBUG(dbgs() << "legalizeIntrinsic: " << MI);
auto IntrinsicID = cast<GIntrinsic>(MI).getIntrinsicID();
-
- if (IntrinsicID == Intrinsic::spv_wave_readlane_first) {
- MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
- LLT DstTy = MRI.getType(MI.getOperand(0).getReg());
- if (needsVectorLegalization(DstTy, *ST)) {
- unsigned MaxVectorSize = ST->isShader() ? 4 : 16;
- unsigned NumElts = llvm::bit_floor(
- std::min<unsigned>(DstTy.getNumElements(), MaxVectorSize));
- return Helper.fewerElementsVectorMultiEltType(
- cast<GIntrinsic>(MI), NumElts) == LegalizerHelper::Legalized;
- }
- }
-
switch (IntrinsicID) {
case Intrinsic::spv_bitcast:
return legalizeSpvBitcast(Helper, MI, GR);
diff --git a/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp b/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp
index e990963210f68a..63f6bf5f4483d9 100644
--- a/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp
@@ -187,12 +187,12 @@ static SPIRVTypeInst deduceTypeFromUses(Register Reg, MachineFunction &MF,
ResType = deduceTypeFromPointerOperand(&Use, Reg, GR, MIB);
break;
case TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS:
- case TargetOpcode::G_INTRINSIC_CONVERGENT:
case TargetOpcode::G_INTRINSIC: {
auto IntrinsicID = cast<GIntrinsic>(Use).getIntrinsicID();
- if (IntrinsicID == Intrinsic::spv_wave_readlane_first ||
- IntrinsicID == Intrinsic::spv_insertelt ||
- IntrinsicID == Intrinsic::spv_extractelt) {
+ if (IntrinsicID == Intrinsic::spv_insertelt) {
+ if (Reg == Use.getOperand(2).getReg())
+ ResType = deduceTypeFromResultRegister(&Use, Reg, GR, MIB);
+ } else if (IntrinsicID == Intrinsic::spv_extractelt) {
if (Reg == Use.getOperand(2).getReg())
ResType = deduceTypeFromResultRegister(&Use, Reg, GR, MIB);
}
@@ -296,13 +296,10 @@ static SPIRVTypeInst deduceResultTypeFromOperands(MachineInstr *I,
case TargetOpcode::G_SHUFFLE_VECTOR:
return deduceTypeFromOperandRange(I, MIB, GR, 1, 3);
case TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS:
- case TargetOpcode::G_INTRINSIC_CONVERGENT:
case TargetOpcode::G_INTRINSIC: {
auto IntrinsicID = cast<GIntrinsic>(I)->getIntrinsicID();
if (IntrinsicID == Intrinsic::spv_gep)
return deduceGEPType(I, GR, MIB);
- if (IntrinsicID == Intrinsic::spv_wave_readlane_first)
- return deduceTypeFromSingleOperand(I, MIB, GR, 2);
break;
}
case TargetOpcode::G_LOAD: {
diff --git a/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
index db23472353e8fb..bb05a15c815402 100644
--- a/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
+++ b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
@@ -1,17 +1,24 @@
-; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv1.5-vulkan-unknown %s -o - | FileCheck %s
-; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv1.5-vulkan-unknown %s -o - -filetype=obj | spirv-val %}
+; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv1.6-unknown-vulkan1.3-compute --spirv-ext=+SPV_EXT_long_vector %s -o - | FileCheck %s
+; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv1.6-unknown-vulkan1.3-compute --spirv-ext=+SPV_EXT_long_vector %s -o - -filetype=obj | spirv-val --target-env vulkan1.3 %}
; Test WaveReadLaneFirst lowering for scalar, vector, and matrix types.
; CHECK: Capability Shader
; CHECK: Capability GroupNonUniformBallot
+; CHECK: Capability LongVectorEXT
+; CHECK: Extension "SPV_EXT_long_vector"
; CHECK-DAG: %[[#uint:]] = OpTypeInt 32 0
; CHECK-DAG: %[[#f32:]] = OpTypeFloat 32
-; CHECK-DAG: %[[#v2_float:]] = OpTypeVector %[[#f32]] 2
; CHECK-DAG: %[[#v4_float:]] = OpTypeVector %[[#f32]] 4
; CHECK-DAG: %[[#bool:]] = OpTypeBool
; CHECK-DAG: %[[#scope:]] = OpConstant %[[#uint]] 3
+; CHECK-DAG: %[[#size5:]] = OpConstant %[[#uint]] 5
+; CHECK-DAG: %[[#v5_float:]] = OpTypeVectorIdEXT %[[#f32]] %[[#size5]]
+; CHECK-DAG: %[[#size6:]] = OpConstant %[[#uint]] 6
+; CHECK-DAG: %[[#v6_float:]] = OpTypeVectorIdEXT %[[#f32]] %[[#size6]]
+; CHECK-DAG: %[[#size12:]] = OpConstant %[[#uint]] 12
+; CHECK-DAG: %[[#v12_float:]] = OpTypeVectorIdEXT %[[#f32]] %[[#size12]]
@wide_f32_5 = internal addrspace(10) global [5 x float] zeroinitializer
@wide_f32_6 = internal addrspace(10) global [6 x float] zeroinitializer
@@ -58,8 +65,7 @@ entry:
define void @test_floatv5() {
entry:
%expr = load <5 x float>, ptr addrspace(10) @wide_f32_5
-; CHECK: OpGroupNonUniformBroadcastFirst %[[#v4_float]] %[[#scope]]
-; CHECK: OpGroupNonUniformBroadcastFirst %[[#f32]] %[[#scope]]
+; CHECK: OpGroupNonUniformBroadcastFirst %[[#v5_float]] %[[#scope]]
%result = call <5 x float> @llvm.spv.wave.readlane.first.v5f32(
<5 x float> %expr)
store <5 x float> %result, ptr addrspace(10) @wide_f32_5
@@ -70,8 +76,7 @@ entry:
define void @test_float2x3() {
entry:
%expr = load <6 x float>, ptr addrspace(10) @wide_f32_6
-; CHECK: OpGroupNonUniformBroadcastFirst %[[#v4_float]] %[[#scope]]
-; CHECK: OpGroupNonUniformBroadcastFirst %[[#v2_float]] %[[#scope]]
+; CHECK: OpGroupNonUniformBroadcastFirst %[[#v6_float]] %[[#scope]]
%result = call <6 x float> @llvm.spv.wave.readlane.first.v6f32(
<6 x float> %expr)
store <6 x float> %result, ptr addrspace(10) @wide_f32_6
@@ -82,7 +87,7 @@ entry:
define void @test_float3x4() {
entry:
%expr = load <12 x float>, ptr addrspace(10) @wide_f32_12
-; CHECK-COUNT-3: OpGroupNonUniformBroadcastFirst %[[#v4_float]] %[[#scope]]
+; CHECK: OpGroupNonUniformBroadcastFirst %[[#v12_float]] %[[#scope]]
%result = call <12 x float> @llvm.spv.wave.readlane.first.v12f32(
<12 x float> %expr)
store <12 x float> %result, ptr addrspace(10) @wide_f32_12
>From 83eccf5bb7e7b8bcd84aae5ea230924dc8fa6f90 Mon Sep 17 00:00:00 2001
From: Joshua Batista <jbatista at microsoft.com>
Date: Mon, 21 Sep 2026 10:30:45 -0700
Subject: [PATCH 6/7] fix failing test
---
.../hlsl-intrinsics/WaveReadLaneFirst.ll | 20 ++++++++++++-------
1 file changed, 13 insertions(+), 7 deletions(-)
diff --git a/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
index bb05a15c815402..811b24605fb574 100644
--- a/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
+++ b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
@@ -26,7 +26,7 @@
; CHECK-LABEL: Begin function test_float
; CHECK: %[[#fexpr:]] = OpFunctionParameter %[[#f32]]
-define float @test_float(float %fexpr) {
+define internal float @test_float(float %fexpr) {
entry:
; CHECK: %[[#]] = OpGroupNonUniformBroadcastFirst %[[#f32]] %[[#scope]] %[[#fexpr]]
%0 = call float @llvm.spv.wave.readlane.first.f32(float %fexpr)
@@ -35,7 +35,7 @@ entry:
; CHECK-LABEL: Begin function test_int
; CHECK: %[[#iexpr:]] = OpFunctionParameter %[[#uint]]
-define i32 @test_int(i32 %iexpr) {
+define internal i32 @test_int(i32 %iexpr) {
entry:
; CHECK: %[[#]] = OpGroupNonUniformBroadcastFirst %[[#uint]] %[[#scope]] %[[#iexpr]]
%0 = call i32 @llvm.spv.wave.readlane.first.i32(i32 %iexpr)
@@ -44,7 +44,7 @@ entry:
; CHECK-LABEL: Begin function test_bool
; CHECK: %[[#bexpr:]] = OpFunctionParameter %[[#bool]]
-define i1 @test_bool(i1 %bexpr) {
+define internal i1 @test_bool(i1 %bexpr) {
entry:
; CHECK: %[[#]] = OpGroupNonUniformBroadcastFirst %[[#bool]] %[[#scope]] %[[#bexpr]]
%0 = call i1 @llvm.spv.wave.readlane.first.i1(i1 %bexpr)
@@ -53,7 +53,7 @@ entry:
; CHECK-LABEL: Begin function test_vfloat
; CHECK: %[[#vfexpr:]] = OpFunctionParameter %[[#v4_float]]
-define <4 x float> @test_vfloat(<4 x float> %vfexpr) {
+define internal <4 x float> @test_vfloat(<4 x float> %vfexpr) {
entry:
; CHECK: %[[#]] = OpGroupNonUniformBroadcastFirst %[[#v4_float]] %[[#scope]] %[[#vfexpr]]
%0 = call <4 x float> @llvm.spv.wave.readlane.first.v4f32(
@@ -62,7 +62,7 @@ entry:
}
; CHECK-LABEL: Begin function test_floatv5
-define void @test_floatv5() {
+define internal void @test_floatv5() {
entry:
%expr = load <5 x float>, ptr addrspace(10) @wide_f32_5
; CHECK: OpGroupNonUniformBroadcastFirst %[[#v5_float]] %[[#scope]]
@@ -73,7 +73,7 @@ entry:
}
; CHECK-LABEL: Begin function test_float2x3
-define void @test_float2x3() {
+define internal void @test_float2x3() {
entry:
%expr = load <6 x float>, ptr addrspace(10) @wide_f32_6
; CHECK: OpGroupNonUniformBroadcastFirst %[[#v6_float]] %[[#scope]]
@@ -84,7 +84,7 @@ entry:
}
; CHECK-LABEL: Begin function test_float3x4
-define void @test_float3x4() {
+define internal void @test_float3x4() {
entry:
%expr = load <12 x float>, ptr addrspace(10) @wide_f32_12
; CHECK: OpGroupNonUniformBroadcastFirst %[[#v12_float]] %[[#scope]]
@@ -94,6 +94,10 @@ entry:
ret void
}
+define void @main() #0 {
+ ret void
+}
+
declare float @llvm.spv.wave.readlane.first.f32(float)
declare i32 @llvm.spv.wave.readlane.first.i32(i32)
declare i1 @llvm.spv.wave.readlane.first.i1(i1)
@@ -101,3 +105,5 @@ declare <4 x float> @llvm.spv.wave.readlane.first.v4f32(<4 x float>)
declare <5 x float> @llvm.spv.wave.readlane.first.v5f32(<5 x float>)
declare <6 x float> @llvm.spv.wave.readlane.first.v6f32(<6 x float>)
declare <12 x float> @llvm.spv.wave.readlane.first.v12f32(<12 x float>)
+
+attributes #0 = { "hlsl.numthreads"="1,1,1" "hlsl.shader"="compute" }
>From 94294487ca0d413971fe14dbc24f4119e02ead3e Mon Sep 17 00:00:00 2001
From: Joshua Batista <jbatista at microsoft.com>
Date: Mon, 21 Sep 2026 17:47:22 -0700
Subject: [PATCH 7/7] split matrix and long-vector tests, and legalize spirv
matrices without depending on SPV_EXT_long_vector
---
.../CodeGen/GlobalISel/LegalizerHelper.cpp | 20 +++++--
llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp | 13 +++++
llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp | 11 ++--
.../DirectX/LongVector/wave-readlane-first.ll | 13 +++++
.../test/CodeGen/DirectX/WaveReadLaneFirst.ll | 33 ------------
.../CodeGen/DirectX/WaveReadLaneFirst_mat.ll | 26 +++++++++
.../wave-readlane-first.ll | 36 +++++++++++++
.../hlsl-intrinsics/WaveReadLaneFirst.ll | 54 ++-----------------
.../hlsl-intrinsics/WaveReadLaneFirst_mat.ll | 50 +++++++++++++++++
9 files changed, 165 insertions(+), 91 deletions(-)
create mode 100644 llvm/test/CodeGen/DirectX/LongVector/wave-readlane-first.ll
create mode 100644 llvm/test/CodeGen/DirectX/WaveReadLaneFirst_mat.ll
create mode 100644 llvm/test/CodeGen/SPIRV/extensions/SPV_EXT_long_vector/wave-readlane-first.ll
create mode 100644 llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst_mat.ll
diff --git a/llvm/lib/CodeGen/GlobalISel/LegalizerHelper.cpp b/llvm/lib/CodeGen/GlobalISel/LegalizerHelper.cpp
index 4ab02b0b12bff5..9f7cd14d4fcdf9 100644
--- a/llvm/lib/CodeGen/GlobalISel/LegalizerHelper.cpp
+++ b/llvm/lib/CodeGen/GlobalISel/LegalizerHelper.cpp
@@ -5230,8 +5230,11 @@ static bool hasSameNumEltsOnAllVectorOperands(
if (!VecTy.isVector())
return false;
unsigned NumElts = VecTy.getNumElements();
+ unsigned IntrinsicIDOp = isa<GIntrinsic>(MI) ? MI.getNumExplicitDefs() : ~0U;
for (unsigned OpIdx = 1; OpIdx < MI.getNumOperands(); ++OpIdx) {
+ if (OpIdx == IntrinsicIDOp)
+ continue;
MachineOperand &Op = MI.getOperand(OpIdx);
if (!Op.isReg()) {
if (!is_contained(NonVecOpIndices, OpIdx))
@@ -5314,8 +5317,10 @@ LegalizerHelper::fewerElementsVectorMultiEltType(
"Non-compatible opcode or not specified non-vector operands");
unsigned OrigNumElts = MRI.getType(MI.getReg(0)).getNumElements();
- unsigned NumInputs = MI.getNumOperands() - MI.getNumDefs();
unsigned NumDefs = MI.getNumDefs();
+ auto *GI = dyn_cast<GIntrinsic>(&MI);
+ unsigned FirstUse = NumDefs + (GI ? 1 : 0);
+ unsigned NumInputs = MI.getNumOperands() - FirstUse;
// Create DstOps (sub-vectors with NumElts elts + Leftover) for each output.
// Build instructions with DstOps to use instruction found by CSE directly.
@@ -5332,7 +5337,7 @@ LegalizerHelper::fewerElementsVectorMultiEltType(
// examples: compare predicate in icmp and fcmp (op 1), vector select with i1
// scalar condition (op 1), immediate in sext_inreg (op 2).
SmallVector<SmallVector<SrcOp, 8>, 3> InputOpsPieces(NumInputs);
- for (unsigned UseIdx = NumDefs, UseNo = 0; UseIdx < MI.getNumOperands();
+ for (unsigned UseIdx = FirstUse, UseNo = 0; UseIdx < MI.getNumOperands();
++UseIdx, ++UseNo) {
if (is_contained(NonVecOpIndices, UseIdx)) {
broadcastSrcOp(InputOpsPieces[UseNo], OutputOpsPieces[0].size(),
@@ -5358,7 +5363,16 @@ LegalizerHelper::fewerElementsVectorMultiEltType(
for (unsigned InputNo = 0; InputNo < NumInputs; ++InputNo)
Uses.push_back(InputOpsPieces[InputNo][i]);
- auto I = MIRBuilder.buildInstr(MI.getOpcode(), Defs, Uses, MI.getFlags());
+ MachineInstrBuilder I;
+ if (GI) {
+ I = MIRBuilder.buildIntrinsic(GI->getIntrinsicID(), Defs,
+ GI->hasSideEffects(), GI->isConvergent());
+ I.setMIFlags(MI.getFlags());
+ for (SrcOp &Use : Uses)
+ Use.addSrcToMIB(I);
+ } else {
+ I = MIRBuilder.buildInstr(MI.getOpcode(), Defs, Uses, MI.getFlags());
+ }
for (unsigned DstNo = 0; DstNo < NumDefs; ++DstNo)
OutputRegs[DstNo].push_back(I.getReg(DstNo));
}
diff --git a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
index dcd89eb3a69d54..484e18fb9cada4 100644
--- a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
@@ -1040,6 +1040,19 @@ bool SPIRVLegalizerInfo::legalizeIntrinsic(LegalizerHelper &Helper,
MachineInstr &MI) const {
LLVM_DEBUG(dbgs() << "legalizeIntrinsic: " << MI);
auto IntrinsicID = cast<GIntrinsic>(MI).getIntrinsicID();
+
+ if (IntrinsicID == Intrinsic::spv_wave_readlane_first) {
+ MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
+ LLT DstTy = MRI.getType(MI.getOperand(0).getReg());
+ if (needsVectorLegalization(DstTy, *ST)) {
+ unsigned MaxVectorSize = ST->isShader() ? 4 : 16;
+ unsigned NumElts = llvm::bit_floor(
+ std::min<unsigned>(DstTy.getNumElements(), MaxVectorSize));
+ return Helper.fewerElementsVectorMultiEltType(
+ cast<GIntrinsic>(MI), NumElts) == LegalizerHelper::Legalized;
+ }
+ }
+
switch (IntrinsicID) {
case Intrinsic::spv_bitcast:
return legalizeSpvBitcast(Helper, MI, GR);
diff --git a/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp b/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp
index 63f6bf5f4483d9..e990963210f68a 100644
--- a/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp
@@ -187,12 +187,12 @@ static SPIRVTypeInst deduceTypeFromUses(Register Reg, MachineFunction &MF,
ResType = deduceTypeFromPointerOperand(&Use, Reg, GR, MIB);
break;
case TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS:
+ case TargetOpcode::G_INTRINSIC_CONVERGENT:
case TargetOpcode::G_INTRINSIC: {
auto IntrinsicID = cast<GIntrinsic>(Use).getIntrinsicID();
- if (IntrinsicID == Intrinsic::spv_insertelt) {
- if (Reg == Use.getOperand(2).getReg())
- ResType = deduceTypeFromResultRegister(&Use, Reg, GR, MIB);
- } else if (IntrinsicID == Intrinsic::spv_extractelt) {
+ if (IntrinsicID == Intrinsic::spv_wave_readlane_first ||
+ IntrinsicID == Intrinsic::spv_insertelt ||
+ IntrinsicID == Intrinsic::spv_extractelt) {
if (Reg == Use.getOperand(2).getReg())
ResType = deduceTypeFromResultRegister(&Use, Reg, GR, MIB);
}
@@ -296,10 +296,13 @@ static SPIRVTypeInst deduceResultTypeFromOperands(MachineInstr *I,
case TargetOpcode::G_SHUFFLE_VECTOR:
return deduceTypeFromOperandRange(I, MIB, GR, 1, 3);
case TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS:
+ case TargetOpcode::G_INTRINSIC_CONVERGENT:
case TargetOpcode::G_INTRINSIC: {
auto IntrinsicID = cast<GIntrinsic>(I)->getIntrinsicID();
if (IntrinsicID == Intrinsic::spv_gep)
return deduceGEPType(I, GR, MIB);
+ if (IntrinsicID == Intrinsic::spv_wave_readlane_first)
+ return deduceTypeFromSingleOperand(I, MIB, GR, 2);
break;
}
case TargetOpcode::G_LOAD: {
diff --git a/llvm/test/CodeGen/DirectX/LongVector/wave-readlane-first.ll b/llvm/test/CodeGen/DirectX/LongVector/wave-readlane-first.ll
new file mode 100644
index 00000000000000..fa55942bce1469
--- /dev/null
+++ b/llvm/test/CodeGen/DirectX/LongVector/wave-readlane-first.ll
@@ -0,0 +1,13 @@
+; RUN: llc -mtriple=dxil-pc-shadermodel6.8-library -o - %s | FileCheck %s --check-prefixes=CHECK,CHECK-SCALAR
+; RUN: llc -mtriple=dxil-pc-shadermodel6.9-library -stop-before=dxil-op-lower -o - %s | FileCheck %s --check-prefixes=CHECK,CHECK-VECTOR
+
+; CHECK-LABEL: define <5 x float> @wave_readlane_first_v5float(
+; CHECK-SCALAR-COUNT-5: call float @dx.op.waveReadLaneFirst.f32(i32 118,
+; CHECK-VECTOR: call <5 x float> @llvm.dx.wave.readlane.first.v5f32
+define <5 x float> @wave_readlane_first_v5float(<5 x float> %expr) {
+ %ret = call <5 x float> @llvm.dx.wave.readlane.first.v5f32(
+ <5 x float> %expr)
+ ret <5 x float> %ret
+}
+
+declare <5 x float> @llvm.dx.wave.readlane.first.v5f32(<5 x float>)
diff --git a/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll b/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll
index c890eef2533472..0a68455aa166e1 100644
--- a/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll
+++ b/llvm/test/CodeGen/DirectX/WaveReadLaneFirst.ll
@@ -84,36 +84,6 @@ entry:
ret <4 x float> %ret
}
-define noundef <5 x float> @wave_readlane_first_v5float(
- <5 x float> noundef %expr) {
-entry:
-; CHECK-LABEL: define noundef <5 x float> @wave_readlane_first_v5float(
-; CHECK-COUNT-5: call float @dx.op.waveReadLaneFirst.f32(i32 118,
- %ret = call <5 x float> @llvm.dx.wave.readlane.first.v5f32(
- <5 x float> %expr)
- ret <5 x float> %ret
-}
-
-define noundef <6 x float> @wave_readlane_first_float2x3(
- <6 x float> noundef %expr) {
-entry:
-; CHECK-LABEL: define noundef <6 x float> @wave_readlane_first_float2x3(
-; CHECK-COUNT-6: call float @dx.op.waveReadLaneFirst.f32(i32 118,
- %ret = call <6 x float> @llvm.dx.wave.readlane.first.v6f32(
- <6 x float> %expr)
- ret <6 x float> %ret
-}
-
-define noundef <12 x float> @wave_readlane_first_float3x4(
- <12 x float> noundef %expr) {
-entry:
-; CHECK-LABEL: define noundef <12 x float> @wave_readlane_first_float3x4(
-; CHECK-COUNT-12: call float @dx.op.waveReadLaneFirst.f32(i32 118,
- %ret = call <12 x float> @llvm.dx.wave.readlane.first.v12f32(
- <12 x float> %expr)
- ret <12 x float> %ret
-}
-
declare half @llvm.dx.wave.readlane.first.f16(half)
declare float @llvm.dx.wave.readlane.first.f32(float)
declare double @llvm.dx.wave.readlane.first.f64(double)
@@ -124,6 +94,3 @@ declare i64 @llvm.dx.wave.readlane.first.i64(i64)
declare <2 x half> @llvm.dx.wave.readlane.first.v2f16(<2 x half>)
declare <3 x i32> @llvm.dx.wave.readlane.first.v3i32(<3 x i32>)
declare <4 x float> @llvm.dx.wave.readlane.first.v4f32(<4 x float>)
-declare <5 x float> @llvm.dx.wave.readlane.first.v5f32(<5 x float>)
-declare <6 x float> @llvm.dx.wave.readlane.first.v6f32(<6 x float>)
-declare <12 x float> @llvm.dx.wave.readlane.first.v12f32(<12 x float>)
diff --git a/llvm/test/CodeGen/DirectX/WaveReadLaneFirst_mat.ll b/llvm/test/CodeGen/DirectX/WaveReadLaneFirst_mat.ll
new file mode 100644
index 00000000000000..cde610839c2915
--- /dev/null
+++ b/llvm/test/CodeGen/DirectX/WaveReadLaneFirst_mat.ll
@@ -0,0 +1,26 @@
+; RUN: opt -S -scalarizer -dxil-op-lower -mtriple=dxil-pc-shadermodel6.3-compute %s | FileCheck %s
+
+; Test WaveReadLaneFirst scalarization for matrix values.
+
+define noundef <6 x float> @wave_readlane_first_float2x3(
+ <6 x float> noundef %expr) {
+entry:
+; CHECK-LABEL: define noundef <6 x float> @wave_readlane_first_float2x3(
+; CHECK-COUNT-6: call float @dx.op.waveReadLaneFirst.f32(i32 118,
+ %ret = call <6 x float> @llvm.dx.wave.readlane.first.v6f32(
+ <6 x float> %expr)
+ ret <6 x float> %ret
+}
+
+define noundef <12 x float> @wave_readlane_first_float3x4(
+ <12 x float> noundef %expr) {
+entry:
+; CHECK-LABEL: define noundef <12 x float> @wave_readlane_first_float3x4(
+; CHECK-COUNT-12: call float @dx.op.waveReadLaneFirst.f32(i32 118,
+ %ret = call <12 x float> @llvm.dx.wave.readlane.first.v12f32(
+ <12 x float> %expr)
+ ret <12 x float> %ret
+}
+
+declare <6 x float> @llvm.dx.wave.readlane.first.v6f32(<6 x float>)
+declare <12 x float> @llvm.dx.wave.readlane.first.v12f32(<12 x float>)
diff --git a/llvm/test/CodeGen/SPIRV/extensions/SPV_EXT_long_vector/wave-readlane-first.ll b/llvm/test/CodeGen/SPIRV/extensions/SPV_EXT_long_vector/wave-readlane-first.ll
new file mode 100644
index 00000000000000..63b00fdbed9884
--- /dev/null
+++ b/llvm/test/CodeGen/SPIRV/extensions/SPV_EXT_long_vector/wave-readlane-first.ll
@@ -0,0 +1,36 @@
+; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv1.6-unknown-vulkan1.3-compute --spirv-ext=+SPV_EXT_long_vector %s -o - | FileCheck %s
+; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv1.6-unknown-vulkan1.3-compute --spirv-ext=+SPV_EXT_long_vector %s -o - -filetype=obj | spirv-val --target-env vulkan1.3 %}
+
+; Test that SPV_EXT_long_vector preserves a WaveReadLaneFirst long vector.
+
+; CHECK-DAG: Capability Shader
+; CHECK-DAG: Capability GroupNonUniformBallot
+; CHECK-DAG: Capability LongVectorEXT
+; CHECK-DAG: Extension "SPV_EXT_long_vector"
+
+; CHECK-DAG: %[[#uint:]] = OpTypeInt 32 0
+; CHECK-DAG: %[[#f32:]] = OpTypeFloat 32
+; CHECK-DAG: %[[#scope:]] = OpConstant %[[#uint]] 3
+; CHECK-DAG: %[[#size5:]] = OpConstant %[[#uint]] 5
+; CHECK-DAG: %[[#v5_float:]] = OpTypeVectorIdEXT %[[#f32]] %[[#size5]]
+
+ at wide_f32_5 = internal addrspace(10) global [5 x float] zeroinitializer
+
+; CHECK-LABEL: Begin function test_floatv5
+define internal void @test_floatv5() {
+entry:
+ %expr = load <5 x float>, ptr addrspace(10) @wide_f32_5
+; CHECK: OpGroupNonUniformBroadcastFirst %[[#v5_float]] %[[#scope]]
+ %result = call <5 x float> @llvm.spv.wave.readlane.first.v5f32(
+ <5 x float> %expr)
+ store <5 x float> %result, ptr addrspace(10) @wide_f32_5
+ ret void
+}
+
+define void @main() #0 {
+ ret void
+}
+
+declare <5 x float> @llvm.spv.wave.readlane.first.v5f32(<5 x float>)
+
+attributes #0 = { "hlsl.numthreads"="1,1,1" "hlsl.shader"="compute" }
diff --git a/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
index 811b24605fb574..e6b111d8f197f8 100644
--- a/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
+++ b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst.ll
@@ -1,28 +1,16 @@
-; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv1.6-unknown-vulkan1.3-compute --spirv-ext=+SPV_EXT_long_vector %s -o - | FileCheck %s
-; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv1.6-unknown-vulkan1.3-compute --spirv-ext=+SPV_EXT_long_vector %s -o - -filetype=obj | spirv-val --target-env vulkan1.3 %}
+; 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 %}
-; Test WaveReadLaneFirst lowering for scalar, vector, and matrix types.
+; Test WaveReadLaneFirst lowering for scalar and vector types.
; CHECK: Capability Shader
; CHECK: Capability GroupNonUniformBallot
-; CHECK: Capability LongVectorEXT
-; CHECK: Extension "SPV_EXT_long_vector"
; CHECK-DAG: %[[#uint:]] = OpTypeInt 32 0
; CHECK-DAG: %[[#f32:]] = OpTypeFloat 32
; CHECK-DAG: %[[#v4_float:]] = OpTypeVector %[[#f32]] 4
; CHECK-DAG: %[[#bool:]] = OpTypeBool
; CHECK-DAG: %[[#scope:]] = OpConstant %[[#uint]] 3
-; CHECK-DAG: %[[#size5:]] = OpConstant %[[#uint]] 5
-; CHECK-DAG: %[[#v5_float:]] = OpTypeVectorIdEXT %[[#f32]] %[[#size5]]
-; CHECK-DAG: %[[#size6:]] = OpConstant %[[#uint]] 6
-; CHECK-DAG: %[[#v6_float:]] = OpTypeVectorIdEXT %[[#f32]] %[[#size6]]
-; CHECK-DAG: %[[#size12:]] = OpConstant %[[#uint]] 12
-; CHECK-DAG: %[[#v12_float:]] = OpTypeVectorIdEXT %[[#f32]] %[[#size12]]
-
- at wide_f32_5 = internal addrspace(10) global [5 x float] zeroinitializer
- at wide_f32_6 = internal addrspace(10) global [6 x float] zeroinitializer
- at wide_f32_12 = internal addrspace(10) global [12 x float] zeroinitializer
; CHECK-LABEL: Begin function test_float
; CHECK: %[[#fexpr:]] = OpFunctionParameter %[[#f32]]
@@ -61,39 +49,6 @@ entry:
ret <4 x float> %0
}
-; CHECK-LABEL: Begin function test_floatv5
-define internal void @test_floatv5() {
-entry:
- %expr = load <5 x float>, ptr addrspace(10) @wide_f32_5
-; CHECK: OpGroupNonUniformBroadcastFirst %[[#v5_float]] %[[#scope]]
- %result = call <5 x float> @llvm.spv.wave.readlane.first.v5f32(
- <5 x float> %expr)
- store <5 x float> %result, ptr addrspace(10) @wide_f32_5
- ret void
-}
-
-; CHECK-LABEL: Begin function test_float2x3
-define internal void @test_float2x3() {
-entry:
- %expr = load <6 x float>, ptr addrspace(10) @wide_f32_6
-; CHECK: OpGroupNonUniformBroadcastFirst %[[#v6_float]] %[[#scope]]
- %result = call <6 x float> @llvm.spv.wave.readlane.first.v6f32(
- <6 x float> %expr)
- store <6 x float> %result, ptr addrspace(10) @wide_f32_6
- ret void
-}
-
-; CHECK-LABEL: Begin function test_float3x4
-define internal void @test_float3x4() {
-entry:
- %expr = load <12 x float>, ptr addrspace(10) @wide_f32_12
-; CHECK: OpGroupNonUniformBroadcastFirst %[[#v12_float]] %[[#scope]]
- %result = call <12 x float> @llvm.spv.wave.readlane.first.v12f32(
- <12 x float> %expr)
- store <12 x float> %result, ptr addrspace(10) @wide_f32_12
- ret void
-}
-
define void @main() #0 {
ret void
}
@@ -102,8 +57,5 @@ declare float @llvm.spv.wave.readlane.first.f32(float)
declare i32 @llvm.spv.wave.readlane.first.i32(i32)
declare i1 @llvm.spv.wave.readlane.first.i1(i1)
declare <4 x float> @llvm.spv.wave.readlane.first.v4f32(<4 x float>)
-declare <5 x float> @llvm.spv.wave.readlane.first.v5f32(<5 x float>)
-declare <6 x float> @llvm.spv.wave.readlane.first.v6f32(<6 x float>)
-declare <12 x float> @llvm.spv.wave.readlane.first.v12f32(<12 x float>)
attributes #0 = { "hlsl.numthreads"="1,1,1" "hlsl.shader"="compute" }
diff --git a/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst_mat.ll b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst_mat.ll
new file mode 100644
index 00000000000000..e1aa0ecfd7aed6
--- /dev/null
+++ b/llvm/test/CodeGen/SPIRV/hlsl-intrinsics/WaveReadLaneFirst_mat.ll
@@ -0,0 +1,50 @@
+; 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 %}
+
+; Test WaveReadLaneFirst lowering for matrix types without long vectors.
+
+; CHECK: Capability Shader
+; CHECK: Capability GroupNonUniformBallot
+; CHECK-NOT: Capability LongVectorEXT
+; CHECK-NOT: Extension "SPV_EXT_long_vector"
+
+; CHECK-DAG: %[[#uint:]] = OpTypeInt 32 0
+; CHECK-DAG: %[[#f32:]] = OpTypeFloat 32
+; CHECK-DAG: %[[#v2_float:]] = OpTypeVector %[[#f32]] 2
+; CHECK-DAG: %[[#v4_float:]] = OpTypeVector %[[#f32]] 4
+; CHECK-DAG: %[[#scope:]] = OpConstant %[[#uint]] 3
+
+ at wide_f32_6 = internal addrspace(10) global [6 x float] zeroinitializer
+ at wide_f32_12 = internal addrspace(10) global [12 x float] zeroinitializer
+
+; CHECK-LABEL: Begin function test_float2x3
+define internal void @test_float2x3() {
+entry:
+ %expr = load <6 x float>, ptr addrspace(10) @wide_f32_6
+; CHECK: OpGroupNonUniformBroadcastFirst %[[#v4_float]] %[[#scope]]
+; CHECK: OpGroupNonUniformBroadcastFirst %[[#v2_float]] %[[#scope]]
+ %result = call <6 x float> @llvm.spv.wave.readlane.first.v6f32(
+ <6 x float> %expr)
+ store <6 x float> %result, ptr addrspace(10) @wide_f32_6
+ ret void
+}
+
+; CHECK-LABEL: Begin function test_float3x4
+define internal void @test_float3x4() {
+entry:
+ %expr = load <12 x float>, ptr addrspace(10) @wide_f32_12
+; CHECK-COUNT-3: OpGroupNonUniformBroadcastFirst %[[#v4_float]] %[[#scope]]
+ %result = call <12 x float> @llvm.spv.wave.readlane.first.v12f32(
+ <12 x float> %expr)
+ store <12 x float> %result, ptr addrspace(10) @wide_f32_12
+ ret void
+}
+
+define void @main() #0 {
+ ret void
+}
+
+declare <6 x float> @llvm.spv.wave.readlane.first.v6f32(<6 x float>)
+declare <12 x float> @llvm.spv.wave.readlane.first.v12f32(<12 x float>)
+
+attributes #0 = { "hlsl.numthreads"="1,1,1" "hlsl.shader"="compute" }
More information about the llvm-commits
mailing list