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 <[email protected]> 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@-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@-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@-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@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@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@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@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 <[email protected]> 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@-1 {{invalid operand of type 'S' where a scalar or vector is required}} + // expected-error@-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@-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 <[email protected]> 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@-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 <[email protected]> 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 +@wide_f32_5 = internal addrspace(10) global [5 x float] zeroinitializer +@wide_f32_6 = internal addrspace(10) global [6 x float] zeroinitializer +@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 <[email protected]> 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 <[email protected]> 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 <[email protected]> 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]] + +@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]] - -@wide_f32_5 = internal addrspace(10) global [5 x float] zeroinitializer -@wide_f32_6 = internal addrspace(10) global [6 x float] zeroinitializer -@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 + +@wide_f32_6 = internal addrspace(10) global [6 x float] zeroinitializer +@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" } _______________________________________________ cfe-commits mailing list [email protected] https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits
