[spirv-reader][ir] Add support for OpGroupNonUniform*Max Add support for `OpGroupNonUniformSMax`, `OpGroupNonUniformUMax`, and `OpGroupNonUniformFMax` and convert them to `subgroupMax`. Fixed: 431033063, 431033704, 431033492 Change-Id: I3268c8f8dc05fd31ec00f7ac7991dbb00a5467a3 Reviewed-on: https://dawn-review.googlesource.com/c/dawn/+/252694 Commit-Queue: dan sinclair <dsinclair@chromium.org> Reviewed-by: David Neto <dneto@google.com>
diff --git a/src/tint/lang/spirv/builtin_fn.cc b/src/tint/lang/spirv/builtin_fn.cc index ee0101d..e569bb6 100644 --- a/src/tint/lang/spirv/builtin_fn.cc +++ b/src/tint/lang/spirv/builtin_fn.cc
@@ -256,6 +256,8 @@ return "group_non_uniform_shuffle_up"; case BuiltinFn::kGroupNonUniformSMin: return "group_non_uniform_s_min"; + case BuiltinFn::kGroupNonUniformSMax: + return "group_non_uniform_s_max"; } return "<unknown>"; } @@ -377,6 +379,7 @@ case BuiltinFn::kGroupNonUniformQuadBroadcast: case BuiltinFn::kGroupNonUniformQuadSwap: case BuiltinFn::kGroupNonUniformSMin: + case BuiltinFn::kGroupNonUniformSMax: break; } return core::ir::Instruction::Accesses{};
diff --git a/src/tint/lang/spirv/builtin_fn.cc.tmpl b/src/tint/lang/spirv/builtin_fn.cc.tmpl index f42214a..861e189 100644 --- a/src/tint/lang/spirv/builtin_fn.cc.tmpl +++ b/src/tint/lang/spirv/builtin_fn.cc.tmpl
@@ -144,6 +144,7 @@ case BuiltinFn::kGroupNonUniformQuadBroadcast: case BuiltinFn::kGroupNonUniformQuadSwap: case BuiltinFn::kGroupNonUniformSMin: + case BuiltinFn::kGroupNonUniformSMax: break; } return core::ir::Instruction::Accesses{};
diff --git a/src/tint/lang/spirv/builtin_fn.h b/src/tint/lang/spirv/builtin_fn.h index eb6d960..0d4100a 100644 --- a/src/tint/lang/spirv/builtin_fn.h +++ b/src/tint/lang/spirv/builtin_fn.h
@@ -155,6 +155,7 @@ kGroupNonUniformShuffleDown, kGroupNonUniformShuffleUp, kGroupNonUniformSMin, + kGroupNonUniformSMax, kNone, };
diff --git a/src/tint/lang/spirv/intrinsic/data.cc b/src/tint/lang/spirv/intrinsic/data.cc index aea05b6..157439c 100644 --- a/src/tint/lang/spirv/intrinsic/data.cc +++ b/src/tint/lang/spirv/intrinsic/data.cc
@@ -12134,6 +12134,13 @@ /* num overloads */ 2, /* overloads */ OverloadIndex(296), }, + { + /* [107] */ + /* fn group_non_uniform_s_max[T : iu32](scope: u32, group_operation: u32, value: T) -> T */ + /* fn group_non_uniform_s_max[N : num, T : iu32](scope: u32, group_operation: u32, value: vec<N, T>) -> vec<N, T> */ + /* num overloads */ 2, + /* overloads */ OverloadIndex(296), + }, }; // clang-format on
diff --git a/src/tint/lang/spirv/reader/lower/builtins.cc b/src/tint/lang/spirv/reader/lower/builtins.cc index 322d360..2b7046d 100644 --- a/src/tint/lang/spirv/reader/lower/builtins.cc +++ b/src/tint/lang/spirv/reader/lower/builtins.cc
@@ -243,7 +243,10 @@ GroupNonUniformQuadSwap(builtin); break; case spirv::BuiltinFn::kGroupNonUniformSMin: - GroupNonUniformSMin(builtin); + GroupNonUniformMinMax(builtin, core::BuiltinFn::kSubgroupMin); + break; + case spirv::BuiltinFn::kGroupNonUniformSMax: + GroupNonUniformMinMax(builtin, core::BuiltinFn::kSubgroupMax); break; case spirv::BuiltinFn::kAtomicLoad: case spirv::BuiltinFn::kAtomicStore: @@ -289,7 +292,7 @@ } } - void GroupNonUniformSMin(spirv::ir::BuiltinCall* call) { + void GroupNonUniformMinMax(spirv::ir::BuiltinCall* call, core::BuiltinFn fn) { auto* value = call->Args()[2]; auto* orig_type = call->Result()->Type(); @@ -300,7 +303,7 @@ value = b.Convert(type, value)->Result(); } - value = b.Call(type, core::BuiltinFn::kSubgroupMin, Vector{value})->Result(); + value = b.Call(type, fn, Vector{value})->Result(); if (type != orig_type) { value = b.Convert(call->Result()->Type(), value)->Result();
diff --git a/src/tint/lang/spirv/reader/lower/builtins_test.cc b/src/tint/lang/spirv/reader/lower/builtins_test.cc index 7a42ef7..684703e 100644 --- a/src/tint/lang/spirv/reader/lower/builtins_test.cc +++ b/src/tint/lang/spirv/reader/lower/builtins_test.cc
@@ -10423,5 +10423,137 @@ EXPECT_EQ(expect, str()); } +TEST_F(SpirvReader_BuiltinsTest, NonUniformSMax_Scalar_i32) { + auto* ep = b.ComputeFunction("main"); + + b.Append(ep->Block(), [&] { // + b.Call<spirv::ir::BuiltinCall>(ty.i32(), spirv::BuiltinFn::kGroupNonUniformSMax, 3_u, 0_u, + 1_i); + b.Return(ep); + }); + + auto src = R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:i32 = spirv.group_non_uniform_s_max 3u, 0u, 1i + ret + } +} +)"; + EXPECT_EQ(src, str()); + + Run(Builtins); + + auto expect = R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:i32 = subgroupMax 1i + ret + } +} +)"; + EXPECT_EQ(expect, str()); +} + +TEST_F(SpirvReader_BuiltinsTest, NonUniformSMax_Vector_i32) { + auto* ep = b.ComputeFunction("main"); + + b.Append(ep->Block(), [&] { // + b.Call<spirv::ir::BuiltinCall>(ty.vec3<i32>(), spirv::BuiltinFn::kGroupNonUniformSMax, 3_u, + 0_u, b.Composite(ty.vec3<i32>(), 1_i, 3_i, 1_i)); + b.Return(ep); + }); + + auto src = R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:vec3<i32> = spirv.group_non_uniform_s_max 3u, 0u, vec3<i32>(1i, 3i, 1i) + ret + } +} +)"; + EXPECT_EQ(src, str()); + + Run(Builtins); + + auto expect = R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:vec3<i32> = subgroupMax vec3<i32>(1i, 3i, 1i) + ret + } +} +)"; + EXPECT_EQ(expect, str()); +} + +TEST_F(SpirvReader_BuiltinsTest, NonUniformSMax_Scalar_u32) { + auto* ep = b.ComputeFunction("main"); + + b.Append(ep->Block(), [&] { // + b.Call<spirv::ir::BuiltinCall>(ty.u32(), spirv::BuiltinFn::kGroupNonUniformSMax, 3_u, 0_u, + 1_u); + b.Return(ep); + }); + + auto src = R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:u32 = spirv.group_non_uniform_s_max 3u, 0u, 1u + ret + } +} +)"; + EXPECT_EQ(src, str()); + + Run(Builtins); + + auto expect = R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:i32 = convert 1u + %3:i32 = subgroupMax %2 + %4:u32 = convert %3 + ret + } +} +)"; + EXPECT_EQ(expect, str()); +} + +TEST_F(SpirvReader_BuiltinsTest, NonUniformSMax_Vector_u32) { + auto* ep = b.ComputeFunction("main"); + + b.Append(ep->Block(), [&] { // + b.Call<spirv::ir::BuiltinCall>(ty.vec3<u32>(), spirv::BuiltinFn::kGroupNonUniformSMax, 3_u, + 0_u, b.Composite(ty.vec3<u32>(), 1_u, 3_u, 1_u)); + b.Return(ep); + }); + + auto src = R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:vec3<u32> = spirv.group_non_uniform_s_max 3u, 0u, vec3<u32>(1u, 3u, 1u) + ret + } +} +)"; + EXPECT_EQ(src, str()); + + Run(Builtins); + + auto expect = R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:vec3<i32> = convert vec3<u32>(1u, 3u, 1u) + %3:vec3<i32> = subgroupMax %2 + %4:vec3<u32> = convert %3 + ret + } +} +)"; + EXPECT_EQ(expect, str()); +} + } // namespace } // namespace tint::spirv::reader::lower
diff --git a/src/tint/lang/spirv/reader/parser/builtin_test.cc b/src/tint/lang/spirv/reader/parser/builtin_test.cc index c977d79..0510359 100644 --- a/src/tint/lang/spirv/reader/parser/builtin_test.cc +++ b/src/tint/lang/spirv/reader/parser/builtin_test.cc
@@ -3164,5 +3164,276 @@ SPV_ENV_VULKAN_1_1); } +TEST_F(SpirvParserTest, NonUniformSMax_Scalar_i32) { + EXPECT_IR_SPV(R"( + OpCapability Shader + OpCapability GroupNonUniformArithmetic + OpMemoryModel Logical GLSL450 + OpEntryPoint GLCompute %main "main" + OpExecutionMode %main LocalSize 1 1 1 + OpName %main "main" + %int = OpTypeInt 32 1 + %int_1 = OpConstant %int 1 + %uint = OpTypeInt 32 0 + %uint_3 = OpConstant %uint 3 + %bool = OpTypeBool + %true = OpConstantTrue %bool + %void = OpTypeVoid + %23 = OpTypeFunction %void + %main = OpFunction %void None %23 + %24 = OpLabel + %8 = OpGroupNonUniformSMax %int %uint_3 Reduce %int_1 + OpReturn + OpFunctionEnd +)", + R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:i32 = spirv.group_non_uniform_s_max 3u, 0u, 1i + ret + } +} +)", + SPV_ENV_VULKAN_1_1); +} + +TEST_F(SpirvParserTest, NonUniformSMax_Vector_i32) { + EXPECT_IR_SPV(R"( + OpCapability Shader + OpCapability GroupNonUniformArithmetic + OpMemoryModel Logical GLSL450 + OpEntryPoint GLCompute %main "main" + OpExecutionMode %main LocalSize 1 1 1 + OpName %main "main" + %uint = OpTypeInt 32 0 + %uint_3 = OpConstant %uint 3 + %int = OpTypeInt 32 1 + %int_1 = OpConstant %int 1 + %int_3 = OpConstant %int 3 + %v3int = OpTypeVector %int 3 + %12 = OpConstantComposite %v3int %int_1 %int_3 %int_1 + %bool = OpTypeBool + %true = OpConstantTrue %bool + %void = OpTypeVoid + %23 = OpTypeFunction %void + %main = OpFunction %void None %23 + %24 = OpLabel + %8 = OpGroupNonUniformSMax %v3int %uint_3 Reduce %12 + OpReturn + OpFunctionEnd +)", + R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:vec3<i32> = spirv.group_non_uniform_s_max 3u, 0u, vec3<i32>(1i, 3i, 1i) + ret + } +} +)", + SPV_ENV_VULKAN_1_1); +} + +TEST_F(SpirvParserTest, NonUniformSMax_Scalar_u32) { + EXPECT_IR_SPV(R"( + OpCapability Shader + OpCapability GroupNonUniformArithmetic + OpMemoryModel Logical GLSL450 + OpEntryPoint GLCompute %main "main" + OpExecutionMode %main LocalSize 1 1 1 + OpName %main "main" + %uint = OpTypeInt 32 0 + %uint_1 = OpConstant %uint 1 + %uint_3 = OpConstant %uint 3 + %bool = OpTypeBool + %true = OpConstantTrue %bool + %void = OpTypeVoid + %23 = OpTypeFunction %void + %main = OpFunction %void None %23 + %24 = OpLabel + %8 = OpGroupNonUniformSMax %uint %uint_3 Reduce %uint_1 + OpReturn + OpFunctionEnd +)", + R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:u32 = spirv.group_non_uniform_s_max 3u, 0u, 1u + ret + } +} +)", + SPV_ENV_VULKAN_1_1); +} + +TEST_F(SpirvParserTest, NonUniformSMax_Vector_u32) { + EXPECT_IR_SPV(R"( + OpCapability Shader + OpCapability GroupNonUniformArithmetic + OpMemoryModel Logical GLSL450 + OpEntryPoint GLCompute %main "main" + OpExecutionMode %main LocalSize 1 1 1 + OpName %main "main" + %uint = OpTypeInt 32 0 + %uint_1 = OpConstant %uint 1 + %uint_3 = OpConstant %uint 3 + %v3uint = OpTypeVector %uint 3 + %12 = OpConstantComposite %v3uint %uint_1 %uint_3 %uint_1 + %bool = OpTypeBool + %true = OpConstantTrue %bool + %void = OpTypeVoid + %23 = OpTypeFunction %void + %main = OpFunction %void None %23 + %24 = OpLabel + %8 = OpGroupNonUniformSMax %v3uint %uint_3 Reduce %12 + OpReturn + OpFunctionEnd +)", + R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:vec3<u32> = spirv.group_non_uniform_s_max 3u, 0u, vec3<u32>(1u, 3u, 1u) + ret + } +} +)", + SPV_ENV_VULKAN_1_1); +} + +TEST_F(SpirvParserTest, NonUniformUMax_Scalar) { + EXPECT_IR_SPV(R"( + OpCapability Shader + OpCapability GroupNonUniformArithmetic + OpMemoryModel Logical GLSL450 + OpEntryPoint GLCompute %main "main" + OpExecutionMode %main LocalSize 1 1 1 + OpName %main "main" + %uint = OpTypeInt 32 0 + %uint_1 = OpConstant %uint 1 + %uint_3 = OpConstant %uint 3 + %bool = OpTypeBool + %true = OpConstantTrue %bool + %void = OpTypeVoid + %23 = OpTypeFunction %void + %main = OpFunction %void None %23 + %24 = OpLabel + %8 = OpGroupNonUniformUMax %uint %uint_3 Reduce %uint_1 + OpReturn + OpFunctionEnd +)", + R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:u32 = subgroupMax 1u + ret + } +} +)", + SPV_ENV_VULKAN_1_1); +} + +TEST_F(SpirvParserTest, NonUniformUMax_Vector) { + EXPECT_IR_SPV(R"( + OpCapability Shader + OpCapability GroupNonUniformArithmetic + OpMemoryModel Logical GLSL450 + OpEntryPoint GLCompute %main "main" + OpExecutionMode %main LocalSize 1 1 1 + OpName %main "main" + %uint = OpTypeInt 32 0 + %uint_1 = OpConstant %uint 1 + %uint_3 = OpConstant %uint 3 + %v3uint = OpTypeVector %uint 3 + %12 = OpConstantComposite %v3uint %uint_1 %uint_3 %uint_1 + %bool = OpTypeBool + %true = OpConstantTrue %bool + %void = OpTypeVoid + %23 = OpTypeFunction %void + %main = OpFunction %void None %23 + %24 = OpLabel + %8 = OpGroupNonUniformUMax %v3uint %uint_3 Reduce %12 + OpReturn + OpFunctionEnd +)", + R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:vec3<u32> = subgroupMax vec3<u32>(1u, 3u, 1u) + ret + } +} +)", + SPV_ENV_VULKAN_1_1); +} + +TEST_F(SpirvParserTest, NonUniformFMax_Scalar) { + EXPECT_IR_SPV(R"( + OpCapability Shader + OpCapability GroupNonUniformArithmetic + OpMemoryModel Logical GLSL450 + OpEntryPoint GLCompute %main "main" + OpExecutionMode %main LocalSize 1 1 1 + OpName %main "main" + %uint = OpTypeInt 32 0 + %uint_3 = OpConstant %uint 3 + %float = OpTypeFloat 32 + %float_1 = OpConstant %float 1 + %bool = OpTypeBool + %true = OpConstantTrue %bool + %void = OpTypeVoid + %23 = OpTypeFunction %void + %main = OpFunction %void None %23 + %24 = OpLabel + %8 = OpGroupNonUniformFMax %float %uint_3 Reduce %float_1 + OpReturn + OpFunctionEnd +)", + R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:f32 = subgroupMax 1.0f + ret + } +} +)", + SPV_ENV_VULKAN_1_1); +} + +TEST_F(SpirvParserTest, NonUniformFMax_Vector) { + EXPECT_IR_SPV(R"( + OpCapability Shader + OpCapability GroupNonUniformArithmetic + OpMemoryModel Logical GLSL450 + OpEntryPoint GLCompute %main "main" + OpExecutionMode %main LocalSize 1 1 1 + OpName %main "main" + %uint = OpTypeInt 32 0 + %uint_1 = OpConstant %uint 1 + %uint_3 = OpConstant %uint 3 + %float = OpTypeFloat 32 + %float_1 = OpConstant %float 1 + %float_3 = OpConstant %float 3 + %v3float = OpTypeVector %float 3 + %12 = OpConstantComposite %v3float %float_1 %float_3 %float_1 + %bool = OpTypeBool + %true = OpConstantTrue %bool + %void = OpTypeVoid + %23 = OpTypeFunction %void + %main = OpFunction %void None %23 + %24 = OpLabel + %8 = OpGroupNonUniformFMax %v3float %uint_3 Reduce %12 + OpReturn + OpFunctionEnd +)", + R"( +%main = @compute @workgroup_size(1u, 1u, 1u) func():void { + $B1: { + %2:vec3<f32> = subgroupMax vec3<f32>(1.0f, 3.0f, 1.0f) + ret + } +} +)", + SPV_ENV_VULKAN_1_1); +} + } // namespace } // namespace tint::spirv::reader
diff --git a/src/tint/lang/spirv/reader/parser/parser.cc b/src/tint/lang/spirv/reader/parser/parser.cc index f15789d..47a4ffc 100644 --- a/src/tint/lang/spirv/reader/parser/parser.cc +++ b/src/tint/lang/spirv/reader/parser/parser.cc
@@ -2247,11 +2247,18 @@ EmitSubgroupBuiltin(inst, spirv::BuiltinFn::kGroupNonUniformQuadSwap); break; case spv::Op::OpGroupNonUniformSMin: - EmitSubgroupSMin(inst); + EmitSubgroupMinMax(inst, spirv::BuiltinFn::kGroupNonUniformSMin); + break; + case spv::Op::OpGroupNonUniformSMax: + EmitSubgroupMinMax(inst, spirv::BuiltinFn::kGroupNonUniformSMax); break; case spv::Op::OpGroupNonUniformUMin: case spv::Op::OpGroupNonUniformFMin: - EmitSubgroupMin(inst); + EmitSubgroupMinMax(inst, core::BuiltinFn::kSubgroupMin); + break; + case spv::Op::OpGroupNonUniformUMax: + case spv::Op::OpGroupNonUniformFMax: + EmitSubgroupMinMax(inst, core::BuiltinFn::kSubgroupMax); break; default: TINT_UNIMPLEMENTED() @@ -2271,32 +2278,31 @@ } } - void EmitSubgroupMin(spvtools::opt::Instruction& inst) { + void EmitSubgroupMinMax(spvtools::opt::Instruction& inst, core::BuiltinFn fn) { ValidateScope(inst); auto group = inst.GetSingleWordInOperand(1); if (static_cast<spv::GroupOperation>(group) != spv::GroupOperation::Reduce) { - TINT_ICE() << "group operand Reduce required for `Min` instructions"; + TINT_ICE() << "group operand Reduce required for `Min`/`Max` instructions"; } - Emit(b_.Call(Type(inst.type_id()), core::BuiltinFn::kSubgroupMin, Args(inst, 4)), - inst.result_id()); + Emit(b_.Call(Type(inst.type_id()), fn, Args(inst, 4)), inst.result_id()); } - void EmitSubgroupSMin(spvtools::opt::Instruction& inst) { + void EmitSubgroupMinMax(spvtools::opt::Instruction& inst, spirv::BuiltinFn fn) { ValidateScope(inst); auto group = inst.GetSingleWordInOperand(1); if (static_cast<spv::GroupOperation>(group) != spv::GroupOperation::Reduce) { - TINT_ICE() << "group operand Reduce required for `Min` instructions"; + TINT_ICE() << "group operand Reduce required for `Min`/`Max` instructions"; } - Emit(b_.Call<spirv::ir::BuiltinCall>( - Type(inst.type_id()), spirv::BuiltinFn::kGroupNonUniformSMin, - Vector{Value(inst.GetSingleWordInOperand(0)), // - b_.Constant(u32(inst.GetSingleWordInOperand(1))), - Value(inst.GetSingleWordInOperand(2))}), - inst.result_id()); + Emit( + b_.Call<spirv::ir::BuiltinCall>(Type(inst.type_id()), fn, // + Vector{Value(inst.GetSingleWordInOperand(0)), // + b_.Constant(u32(inst.GetSingleWordInOperand(1))), + Value(inst.GetSingleWordInOperand(2))}), + inst.result_id()); } void EmitSubgroupBuiltin(spvtools::opt::Instruction& inst, spirv::BuiltinFn fn) {
diff --git a/src/tint/lang/spirv/spirv.def b/src/tint/lang/spirv/spirv.def index 51e4b29..7892fbf 100644 --- a/src/tint/lang/spirv/spirv.def +++ b/src/tint/lang/spirv/spirv.def
@@ -1861,3 +1861,8 @@ @must_use @stage("fragment", "compute") implicit(N: num, T: iu32) fn group_non_uniform_s_min(scope: u32, group_operation: u32, value: vec<N, T>) -> vec<N, T> +@must_use @stage("fragment", "compute") implicit(T: iu32) +fn group_non_uniform_s_max(scope: u32, group_operation: u32, value: T) -> T +@must_use @stage("fragment", "compute") implicit(N: num, T: iu32) +fn group_non_uniform_s_max(scope: u32, group_operation: u32, value: vec<N, T>) -> vec<N, T> +
diff --git a/src/tint/lang/spirv/writer/printer/printer.cc b/src/tint/lang/spirv/writer/printer/printer.cc index 25b4dd0..05ccdbe 100644 --- a/src/tint/lang/spirv/writer/printer/printer.cc +++ b/src/tint/lang/spirv/writer/printer/printer.cc
@@ -1712,6 +1712,9 @@ case BuiltinFn::kGroupNonUniformSMin: op = spv::Op::OpGroupNonUniformSMin; break; + case BuiltinFn::kGroupNonUniformSMax: + op = spv::Op::OpGroupNonUniformSMax; + break; case spirv::BuiltinFn::kNone: TINT_ICE() << "undefined spirv ir function"; }