opt: fix Subgroup* trimming (#5706)
PR #5648 added support for the GroupNonUniformPartitionedNV. But there
was an issue: the opcodes are enabled by multiple capabilities, and the
actual operand is what matters.
Added testing coverage and the implementation to correctly trim a few
NonUniform capabilities.
Signed-off-by: Nathan Gauër <brioche@google.com>
diff --git a/source/opt/trim_capabilities_pass.cpp b/source/opt/trim_capabilities_pass.cpp
index 24f9e46..aaf4d32 100644
--- a/source/opt/trim_capabilities_pass.cpp
+++ b/source/opt/trim_capabilities_pass.cpp
@@ -399,6 +399,33 @@
return output;
}
+
+bool hasOpcodeConflictingCapabilities(spv::Op opcode) {
+ switch (opcode) {
+ case spv::Op::OpBeginInvocationInterlockEXT:
+ case spv::Op::OpEndInvocationInterlockEXT:
+ case spv::Op::OpGroupNonUniformIAdd:
+ case spv::Op::OpGroupNonUniformFAdd:
+ case spv::Op::OpGroupNonUniformIMul:
+ case spv::Op::OpGroupNonUniformFMul:
+ case spv::Op::OpGroupNonUniformSMin:
+ case spv::Op::OpGroupNonUniformUMin:
+ case spv::Op::OpGroupNonUniformFMin:
+ case spv::Op::OpGroupNonUniformSMax:
+ case spv::Op::OpGroupNonUniformUMax:
+ case spv::Op::OpGroupNonUniformFMax:
+ case spv::Op::OpGroupNonUniformBitwiseAnd:
+ case spv::Op::OpGroupNonUniformBitwiseOr:
+ case spv::Op::OpGroupNonUniformBitwiseXor:
+ case spv::Op::OpGroupNonUniformLogicalAnd:
+ case spv::Op::OpGroupNonUniformLogicalOr:
+ case spv::Op::OpGroupNonUniformLogicalXor:
+ return true;
+ default:
+ return false;
+ }
+}
+
} // namespace
TrimCapabilitiesPass::TrimCapabilitiesPass()
@@ -416,10 +443,7 @@
void TrimCapabilitiesPass::addInstructionRequirementsForOpcode(
spv::Op opcode, CapabilitySet* capabilities,
ExtensionSet* extensions) const {
- // Ignoring OpBeginInvocationInterlockEXT and OpEndInvocationInterlockEXT
- // because they have three possible capabilities, only one of which is needed
- if (opcode == spv::Op::OpBeginInvocationInterlockEXT ||
- opcode == spv::Op::OpEndInvocationInterlockEXT) {
+ if (hasOpcodeConflictingCapabilities(opcode)) {
return;
}
diff --git a/source/opt/trim_capabilities_pass.h b/source/opt/trim_capabilities_pass.h
index 81c07b8..3ff6dba 100644
--- a/source/opt/trim_capabilities_pass.h
+++ b/source/opt/trim_capabilities_pass.h
@@ -81,6 +81,11 @@
spv::Capability::FragmentShaderPixelInterlockEXT,
spv::Capability::FragmentShaderSampleInterlockEXT,
spv::Capability::FragmentShaderShadingRateInterlockEXT,
+ spv::Capability::GroupNonUniform,
+ spv::Capability::GroupNonUniformArithmetic,
+ spv::Capability::GroupNonUniformClustered,
+ spv::Capability::GroupNonUniformPartitionedNV,
+ spv::Capability::GroupNonUniformVote,
spv::Capability::Groups,
spv::Capability::ImageMSArray,
spv::Capability::Int16,
@@ -99,7 +104,6 @@
spv::Capability::StorageUniform16,
spv::Capability::StorageUniformBufferBlock16,
spv::Capability::VulkanMemoryModelDeviceScope,
- spv::Capability::GroupNonUniformPartitionedNV
// clang-format on
};
diff --git a/test/opt/trim_capabilities_pass_test.cpp b/test/opt/trim_capabilities_pass_test.cpp
index d74ccdf..9b2d767 100644
--- a/test/opt/trim_capabilities_pass_test.cpp
+++ b/test/opt/trim_capabilities_pass_test.cpp
@@ -2660,6 +2660,505 @@
EXPECT_EQ(std::get<1>(result), Pass::Status::SuccessWithoutChange);
}
+TEST_F(TrimCapabilitiesPassTest, GroupNonUniform_RemovedWhenUnused) {
+ const std::string kTest = R"(
+ OpCapability Shader
+ OpCapability GroupNonUniformVote
+; CHECK-NOT: OpCapability GroupNonUniformVote
+ OpCapability GroupNonUniformArithmetic
+; CHECK-NOT: OpCapability GroupNonUniformArithmetic
+ OpCapability GroupNonUniformClustered
+; CHECK-NOT: OpCapability GroupNonUniformClustered
+ OpCapability GroupNonUniformPartitionedNV
+; CHECK-NOT: OpCapability GroupNonUniformPartitionedNV
+ OpCapability GroupNonUniform
+; CHECK-NOT: OpCapability GroupNonUniform
+ OpExtension "SPV_NV_shader_subgroup_partitioned"
+; CHECK-NOT: OpExtension "SPV_NV_shader_subgroup_partitioned"
+ OpMemoryModel Logical GLSL450
+ OpEntryPoint GLCompute %main "main"
+ OpExecutionMode %main LocalSize 1 2 4
+ %void = OpTypeVoid
+ %3 = OpTypeFunction %void
+ %main = OpFunction %void None %3
+ %6 = OpLabel
+ OpReturn
+ OpFunctionEnd;
+ )";
+ const auto result = SinglePassRunAndMatch<TrimCapabilitiesPass>(
+ kTest, /* do_validation= */ false);
+ EXPECT_EQ(std::get<1>(result), Pass::Status::SuccessWithChange);
+}
+
+TEST_F(TrimCapabilitiesPassTest,
+ GroupNonUniform_RemainsGroupNonUniformWhenInUse) {
+ const std::string kTest = R"(
+ OpCapability GroupNonUniformVote
+; CHECK-NOT: OpCapability GroupNonUniformVote
+ OpCapability GroupNonUniformArithmetic
+; CHECK-NOT: OpCapability GroupNonUniformArithmetic
+ OpCapability GroupNonUniformClustered
+; CHECK-NOT: OpCapability GroupNonUniformClustered
+ OpCapability GroupNonUniformPartitionedNV
+; CHECK-NOT: OpCapability GroupNonUniformPartitionedNV
+ OpCapability GroupNonUniform
+; CHECK: OpCapability GroupNonUniform
+ OpCapability Shader
+ OpExtension "SPV_NV_shader_subgroup_partitioned"
+; CHECK-NOT: OpExtension "SPV_NV_shader_subgroup_partitioned"
+ OpMemoryModel Logical GLSL450
+ OpEntryPoint GLCompute %main "main"
+ OpExecutionMode %main LocalSize 1 2 4
+ %void = OpTypeVoid
+ %bool = OpTypeBool
+ %uint = OpTypeInt 32 0
+ %scope_subgroup = OpConstant %uint 3
+ %3 = OpTypeFunction %void
+ %main = OpFunction %void None %3
+ %6 = OpLabel
+ %7 = OpGroupNonUniformElect %bool %scope_subgroup
+ OpReturn
+ OpFunctionEnd;
+ )";
+ const auto result = SinglePassRunAndMatch<TrimCapabilitiesPass>(
+ kTest, /* do_validation= */ false);
+ EXPECT_EQ(std::get<1>(result), Pass::Status::SuccessWithChange);
+}
+
+TEST_F(TrimCapabilitiesPassTest,
+ GroupNonUniformVote_Remains_OpGroupNonUniformAll) {
+ const std::string kTest = R"(
+ OpCapability Shader
+ OpCapability GroupNonUniformVote
+; CHECK: OpCapability GroupNonUniformVote
+ OpCapability GroupNonUniformArithmetic
+; CHECK-NOT: OpCapability GroupNonUniformArithmetic
+ OpCapability GroupNonUniformClustered
+; CHECK-NOT: OpCapability GroupNonUniformClustered
+ OpCapability GroupNonUniformPartitionedNV
+; CHECK-NOT: OpCapability GroupNonUniformPartitionedNV
+ OpCapability GroupNonUniform
+; CHECK-NOT: OpCapability GroupNonUniform
+ OpExtension "SPV_NV_shader_subgroup_partitioned"
+; CHECK-NOT: OpExtension "SPV_NV_shader_subgroup_partitioned"
+ OpMemoryModel Logical GLSL450
+ OpEntryPoint GLCompute %main "main"
+ OpExecutionMode %main LocalSize 1 2 4
+ %void = OpTypeVoid
+ %bool = OpTypeBool
+ %uint = OpTypeInt 32 0
+ %scope_subgroup = OpConstant %uint 3
+ %true = OpConstantTrue %bool
+ %3 = OpTypeFunction %void
+ %main = OpFunction %void None %3
+ %6 = OpLabel
+ %7 = OpGroupNonUniformAll %bool %scope_subgroup %true
+ OpReturn
+ OpFunctionEnd;
+ )";
+ const auto result = SinglePassRunAndMatch<TrimCapabilitiesPass>(
+ kTest, /* do_validation= */ false);
+ EXPECT_EQ(std::get<1>(result), Pass::Status::SuccessWithChange);
+}
+
+TEST_F(TrimCapabilitiesPassTest,
+ GroupNonUniformVote_Remains_OpGroupNonUniformAny) {
+ const std::string kTest = R"(
+ OpCapability Shader
+ OpCapability GroupNonUniformVote
+; CHECK: OpCapability GroupNonUniformVote
+ OpCapability GroupNonUniformArithmetic
+; CHECK-NOT: OpCapability GroupNonUniformArithmetic
+ OpCapability GroupNonUniformClustered
+; CHECK-NOT: OpCapability GroupNonUniformClustered
+ OpCapability GroupNonUniformPartitionedNV
+; CHECK-NOT: OpCapability GroupNonUniformPartitionedNV
+ OpCapability GroupNonUniform
+; CHECK-NOT: OpCapability GroupNonUniform
+ OpExtension "SPV_NV_shader_subgroup_partitioned"
+; CHECK-NOT: OpExtension "SPV_NV_shader_subgroup_partitioned"
+ OpMemoryModel Logical GLSL450
+ OpEntryPoint GLCompute %main "main"
+ OpExecutionMode %main LocalSize 1 2 4
+ %void = OpTypeVoid
+ %bool = OpTypeBool
+ %uint = OpTypeInt 32 0
+ %scope_subgroup = OpConstant %uint 3
+ %true = OpConstantTrue %bool
+ %3 = OpTypeFunction %void
+ %main = OpFunction %void None %3
+ %6 = OpLabel
+ %7 = OpGroupNonUniformAny %bool %scope_subgroup %true
+ OpReturn
+ OpFunctionEnd;
+ )";
+ const auto result = SinglePassRunAndMatch<TrimCapabilitiesPass>(
+ kTest, /* do_validation= */ false);
+ EXPECT_EQ(std::get<1>(result), Pass::Status::SuccessWithChange);
+}
+
+TEST_F(TrimCapabilitiesPassTest,
+ GroupNonUniformArithmetic_Remains_OpGroupNonUniformIAdd_Reduce) {
+ const std::string kTest = R"(
+ OpCapability Shader
+ OpCapability GroupNonUniformVote
+; CHECK-NOT: OpCapability GroupNonUniformVote
+ OpCapability GroupNonUniformArithmetic
+; CHECK: OpCapability GroupNonUniformArithmetic
+ OpCapability GroupNonUniformClustered
+; CHECK-NOT: OpCapability GroupNonUniformClustered
+ OpCapability GroupNonUniformPartitionedNV
+; CHECK-NOT: OpCapability GroupNonUniformPartitionedNV
+ OpCapability GroupNonUniform
+; CHECK-NOT: OpCapability GroupNonUniform
+ OpExtension "SPV_NV_shader_subgroup_partitioned"
+; CHECK-NOT: OpExtension "SPV_NV_shader_subgroup_partitioned"
+ OpMemoryModel Logical GLSL450
+ OpEntryPoint GLCompute %main "main"
+ OpExecutionMode %main LocalSize 1 2 4
+ %void = OpTypeVoid
+ %bool = OpTypeBool
+ %uint = OpTypeInt 32 0
+ %scope_subgroup = OpConstant %uint 3
+ %uint_1 = OpConstant %uint 1
+ %true = OpConstantTrue %bool
+ %3 = OpTypeFunction %void
+ %main = OpFunction %void None %3
+ %6 = OpLabel
+ %7 = OpGroupNonUniformIAdd %uint %scope_subgroup Reduce %uint_1
+ OpReturn
+ OpFunctionEnd;
+ )";
+ const auto result = SinglePassRunAndMatch<TrimCapabilitiesPass>(
+ kTest, /* do_validation= */ true);
+ EXPECT_EQ(std::get<1>(result), Pass::Status::SuccessWithChange);
+}
+
+TEST_F(TrimCapabilitiesPassTest,
+ GroupNonUniformArithmetic_Remains_OpGroupNonUniformIAdd_InclusiveScan) {
+ const std::string kTest = R"(
+ OpCapability Shader
+ OpCapability GroupNonUniformVote
+; CHECK-NOT: OpCapability GroupNonUniformVote
+ OpCapability GroupNonUniformArithmetic
+; CHECK: OpCapability GroupNonUniformArithmetic
+ OpCapability GroupNonUniformClustered
+; CHECK-NOT: OpCapability GroupNonUniformClustered
+ OpCapability GroupNonUniformPartitionedNV
+; CHECK-NOT: OpCapability GroupNonUniformPartitionedNV
+ OpCapability GroupNonUniform
+; CHECK-NOT: OpCapability GroupNonUniform
+ OpExtension "SPV_NV_shader_subgroup_partitioned"
+; CHECK-NOT: OpExtension "SPV_NV_shader_subgroup_partitioned"
+ OpMemoryModel Logical GLSL450
+ OpEntryPoint GLCompute %main "main"
+ OpExecutionMode %main LocalSize 1 2 4
+ %void = OpTypeVoid
+ %bool = OpTypeBool
+ %uint = OpTypeInt 32 0
+ %scope_subgroup = OpConstant %uint 3
+ %uint_1 = OpConstant %uint 1
+ %true = OpConstantTrue %bool
+ %3 = OpTypeFunction %void
+ %main = OpFunction %void None %3
+ %6 = OpLabel
+ %7 = OpGroupNonUniformIAdd %uint %scope_subgroup InclusiveScan %uint_1
+ OpReturn
+ OpFunctionEnd;
+ )";
+ const auto result = SinglePassRunAndMatch<TrimCapabilitiesPass>(
+ kTest, /* do_validation= */ true);
+ EXPECT_EQ(std::get<1>(result), Pass::Status::SuccessWithChange);
+}
+
+TEST_F(TrimCapabilitiesPassTest,
+ GroupNonUniformArithmetic_Remains_OpGroupNonUniformIAdd_ExclusiveScan) {
+ const std::string kTest = R"(
+ OpCapability Shader
+ OpCapability GroupNonUniformVote
+; CHECK-NOT: OpCapability GroupNonUniformVote
+ OpCapability GroupNonUniformArithmetic
+; CHECK: OpCapability GroupNonUniformArithmetic
+ OpCapability GroupNonUniformClustered
+; CHECK-NOT: OpCapability GroupNonUniformClustered
+ OpCapability GroupNonUniformPartitionedNV
+; CHECK-NOT: OpCapability GroupNonUniformPartitionedNV
+ OpCapability GroupNonUniform
+; CHECK-NOT: OpCapability GroupNonUniform
+ OpExtension "SPV_NV_shader_subgroup_partitioned"
+; CHECK-NOT: OpExtension "SPV_NV_shader_subgroup_partitioned"
+ OpMemoryModel Logical GLSL450
+ OpEntryPoint GLCompute %main "main"
+ OpExecutionMode %main LocalSize 1 2 4
+ %void = OpTypeVoid
+ %bool = OpTypeBool
+ %uint = OpTypeInt 32 0
+ %scope_subgroup = OpConstant %uint 3
+ %uint_1 = OpConstant %uint 1
+ %true = OpConstantTrue %bool
+ %3 = OpTypeFunction %void
+ %main = OpFunction %void None %3
+ %6 = OpLabel
+ %7 = OpGroupNonUniformIAdd %uint %scope_subgroup ExclusiveScan %uint_1
+ OpReturn
+ OpFunctionEnd;
+ )";
+ const auto result = SinglePassRunAndMatch<TrimCapabilitiesPass>(
+ kTest, /* do_validation= */ true);
+ EXPECT_EQ(std::get<1>(result), Pass::Status::SuccessWithChange);
+}
+
+TEST_F(TrimCapabilitiesPassTest,
+ GroupNonUniformClustered_Remains_OpGroupNonUniformIAdd_ClusteredReduce) {
+ const std::string kTest = R"(
+ OpCapability Shader
+ OpCapability GroupNonUniformVote
+; CHECK-NOT: OpCapability GroupNonUniformVote
+ OpCapability GroupNonUniformArithmetic
+; CHECK-NOT: OpCapability GroupNonUniformArithmetic
+ OpCapability GroupNonUniformClustered
+; CHECK: OpCapability GroupNonUniformClustered
+ OpCapability GroupNonUniformPartitionedNV
+; CHECK-NOT: OpCapability GroupNonUniformPartitionedNV
+ OpCapability GroupNonUniform
+; CHECK-NOT: OpCapability GroupNonUniform
+ OpExtension "SPV_NV_shader_subgroup_partitioned"
+; CHECK-NOT: OpExtension "SPV_NV_shader_subgroup_partitioned"
+ OpMemoryModel Logical GLSL450
+ OpEntryPoint GLCompute %main "main"
+ OpExecutionMode %main LocalSize 1 2 4
+ %void = OpTypeVoid
+ %bool = OpTypeBool
+ %uint = OpTypeInt 32 0
+ %scope_subgroup = OpConstant %uint 3
+ %uint_1 = OpConstant %uint 1
+ %true = OpConstantTrue %bool
+ %3 = OpTypeFunction %void
+ %main = OpFunction %void None %3
+ %6 = OpLabel
+ %7 = OpGroupNonUniformIAdd %uint %scope_subgroup ClusteredReduce %uint_1 %uint_1
+ OpReturn
+ OpFunctionEnd;
+ )";
+ const auto result = SinglePassRunAndMatch<TrimCapabilitiesPass>(
+ kTest, /* do_validation= */ true);
+ EXPECT_EQ(std::get<1>(result), Pass::Status::SuccessWithChange);
+}
+
+struct SubgroupTestCase {
+ // The result type of the subgroup instruction.
+ std::string resultType;
+ // The opcode of the subgroup instruction.
+ std::string opcode;
+ // The actual operand of the subgroup instruction.
+ std::string operand;
+};
+
+static const std::vector<SubgroupTestCase> kSubgroupTestCases{
+ // clang-format off
+ { "uint", "OpGroupNonUniformIAdd", "uint_1" },
+ { "float", "OpGroupNonUniformFAdd", "float_1" },
+ { "uint", "OpGroupNonUniformIMul", "uint_1" },
+ { "float", "OpGroupNonUniformFMul", "float_1" },
+ { "int", "OpGroupNonUniformSMin", "int_1" },
+ { "uint", "OpGroupNonUniformUMin", "uint_1" },
+ { "float", "OpGroupNonUniformFMin", "float_1" },
+ { "int", "OpGroupNonUniformSMax", "int_1" },
+ { "uint", "OpGroupNonUniformUMax", "uint_1" },
+ { "float", "OpGroupNonUniformFMax", "float_1" },
+ { "uint", "OpGroupNonUniformBitwiseAnd", "uint_1" },
+ { "uint", "OpGroupNonUniformBitwiseOr", "uint_1" },
+ { "uint", "OpGroupNonUniformBitwiseXor", "uint_1" },
+ { "bool", "OpGroupNonUniformLogicalAnd", "true" },
+ { "bool", "OpGroupNonUniformLogicalOr", "true" },
+ { "bool", "OpGroupNonUniformLogicalXor", "true" }
+ // clang-format on
+};
+
+using TrimCapabilitiesPassTestSubgroupNV_Unsigned = PassTest<
+ ::testing::TestWithParam<std::tuple<SubgroupTestCase, std::string>>>;
+TEST_P(TrimCapabilitiesPassTestSubgroupNV_Unsigned,
+ GroupNonUniformPartitionedNV_Remains) {
+ SubgroupTestCase test_case = std::get<0>(GetParam());
+ const std::string operation = std::get<1>(GetParam());
+
+ const std::string kTest = R"(
+ OpCapability Shader
+ OpCapability GroupNonUniformVote
+; CHECK-NOT: OpCapability GroupNonUniformVote
+ OpCapability GroupNonUniformArithmetic
+; CHECK-NOT: OpCapability GroupNonUniformArithmetic
+ OpCapability GroupNonUniformClustered
+; CHECK-NOT: OpCapability GroupNonUniformClustered
+ OpCapability GroupNonUniformPartitionedNV
+; CHECK: OpCapability GroupNonUniformPartitionedNV
+ OpCapability GroupNonUniform
+; CHECK-NOT: OpCapability GroupNonUniform
+ OpExtension "SPV_NV_shader_subgroup_partitioned"
+; CHECK: OpExtension "SPV_NV_shader_subgroup_partitioned"
+ OpMemoryModel Logical GLSL450
+ OpEntryPoint GLCompute %main "main"
+ OpExecutionMode %main LocalSize 1 2 4
+ %void = OpTypeVoid
+ %bool = OpTypeBool
+ %uint = OpTypeInt 32 0
+ %int = OpTypeInt 32 1
+ %float = OpTypeFloat 32
+ %v4uint = OpTypeVector %uint 4
+ %scope_subgroup = OpConstant %uint 3
+ %uint_1 = OpConstant %uint 1
+ %int_1 = OpConstant %int 1
+ %float_1 = OpConstant %float 1
+ %uint4_1111 = OpConstantComposite %v4uint %uint_1 %uint_1 %uint_1 %uint_1
+ %true = OpConstantTrue %bool
+ %3 = OpTypeFunction %void
+ %main = OpFunction %void None %3
+ %6 = OpLabel
+ %7 = )" + test_case.opcode +
+ " %" + test_case.resultType + " %scope_subgroup " +
+ operation + " %" + test_case.operand +
+ R"( %uint4_1111
+ OpReturn
+ OpFunctionEnd;
+ )";
+ const auto result = SinglePassRunAndMatch<TrimCapabilitiesPass>(
+ kTest, /* do_validation= */ true);
+ EXPECT_EQ(std::get<1>(result), Pass::Status::SuccessWithChange);
+}
+
+INSTANTIATE_TEST_SUITE_P(
+ TrimCapabilitiesPassTestSubgroupNV_Unsigned_I,
+ TrimCapabilitiesPassTestSubgroupNV_Unsigned,
+ ::testing::Combine(::testing::ValuesIn(kSubgroupTestCases),
+ ::testing::Values("PartitionedReduceNV",
+ "PartitionedInclusiveScanNV",
+ "PartitionedExclusiveScanNV")),
+ [](const ::testing::TestParamInfo<
+ TrimCapabilitiesPassTestSubgroupNV_Unsigned::ParamType>& info) {
+ return std::get<0>(info.param).opcode + "_" + std::get<1>(info.param);
+ });
+
+using TrimCapabilitiesPassTestSubgroupArithmetic_Unsigned = PassTest<
+ ::testing::TestWithParam<std::tuple<SubgroupTestCase, std::string>>>;
+TEST_P(TrimCapabilitiesPassTestSubgroupArithmetic_Unsigned,
+ GroupNonUniformPartitionedArithmetic_Remains) {
+ SubgroupTestCase test_case = std::get<0>(GetParam());
+ const std::string operation = std::get<1>(GetParam());
+
+ const std::string kTest = R"(
+ OpCapability Shader
+ OpCapability GroupNonUniformVote
+; CHECK-NOT: OpCapability GroupNonUniformVote
+ OpCapability GroupNonUniformArithmetic
+; CHECK: OpCapability GroupNonUniformArithmetic
+ OpCapability GroupNonUniformClustered
+; CHECK-NOT: OpCapability GroupNonUniformClustered
+ OpCapability GroupNonUniformPartitionedNV
+; CHECK-NOT: OpCapability GroupNonUniformPartitionedNV
+ OpCapability GroupNonUniform
+; CHECK-NOT: OpCapability GroupNonUniform
+ OpExtension "SPV_NV_shader_subgroup_partitioned"
+; CHECK-NOT: OpExtension "SPV_NV_shader_subgroup_partitioned"
+ OpMemoryModel Logical GLSL450
+ OpEntryPoint GLCompute %main "main"
+ OpExecutionMode %main LocalSize 1 2 4
+ %void = OpTypeVoid
+ %bool = OpTypeBool
+ %uint = OpTypeInt 32 0
+ %int = OpTypeInt 32 1
+ %float = OpTypeFloat 32
+ %v4uint = OpTypeVector %uint 4
+ %scope_subgroup = OpConstant %uint 3
+ %uint_1 = OpConstant %uint 1
+ %int_1 = OpConstant %int 1
+ %float_1 = OpConstant %float 1
+ %uint4_1111 = OpConstantComposite %v4uint %uint_1 %uint_1 %uint_1 %uint_1
+ %true = OpConstantTrue %bool
+ %3 = OpTypeFunction %void
+ %main = OpFunction %void None %3
+ %6 = OpLabel
+ %7 = )" + test_case.opcode +
+ " %" + test_case.resultType + " %scope_subgroup " +
+ operation + " %" + test_case.operand + R"( %uint_1
+ OpReturn
+ OpFunctionEnd;
+ )";
+ const auto result = SinglePassRunAndMatch<TrimCapabilitiesPass>(
+ kTest, /* do_validation= */ true);
+ EXPECT_EQ(std::get<1>(result), Pass::Status::SuccessWithChange);
+}
+
+INSTANTIATE_TEST_SUITE_P(
+ TrimCapabilitiesPassTestSubgroupArithmetic_Unsigned_I,
+ TrimCapabilitiesPassTestSubgroupArithmetic_Unsigned,
+ ::testing::Combine(::testing::ValuesIn(kSubgroupTestCases),
+ ::testing::Values("Reduce", "InclusiveScan",
+ "ExclusiveScan")),
+ [](const ::testing::TestParamInfo<
+ TrimCapabilitiesPassTestSubgroupArithmetic_Unsigned::ParamType>& info) {
+ return std::get<0>(info.param).opcode + "_" + std::get<1>(info.param);
+ });
+
+using TrimCapabilitiesPassTestSubgroupClustered_Unsigned = PassTest<
+ ::testing::TestWithParam<std::tuple<SubgroupTestCase, std::string>>>;
+TEST_P(TrimCapabilitiesPassTestSubgroupClustered_Unsigned,
+ GroupNonUniformPartitionedClustered_Remains) {
+ SubgroupTestCase test_case = std::get<0>(GetParam());
+ const std::string operation = std::get<1>(GetParam());
+
+ const std::string kTest = R"(
+ OpCapability Shader
+ OpCapability GroupNonUniformVote
+; CHECK-NOT: OpCapability GroupNonUniformVote
+ OpCapability GroupNonUniformArithmetic
+; CHECK-NOT: OpCapability GroupNonUniformArithmetic
+ OpCapability GroupNonUniformClustered
+; CHECK: OpCapability GroupNonUniformClustered
+ OpCapability GroupNonUniformPartitionedNV
+; CHECK-NOT: OpCapability GroupNonUniformPartitionedNV
+ OpCapability GroupNonUniform
+; CHECK-NOT: OpCapability GroupNonUniform
+ OpExtension "SPV_NV_shader_subgroup_partitioned"
+; CHECK-NOT: OpExtension "SPV_NV_shader_subgroup_partitioned"
+ OpMemoryModel Logical GLSL450
+ OpEntryPoint GLCompute %main "main"
+ OpExecutionMode %main LocalSize 1 2 4
+ %void = OpTypeVoid
+ %bool = OpTypeBool
+ %uint = OpTypeInt 32 0
+ %int = OpTypeInt 32 1
+ %float = OpTypeFloat 32
+ %v4uint = OpTypeVector %uint 4
+ %scope_subgroup = OpConstant %uint 3
+ %uint_1 = OpConstant %uint 1
+ %int_1 = OpConstant %int 1
+ %float_1 = OpConstant %float 1
+ %uint4_1111 = OpConstantComposite %v4uint %uint_1 %uint_1 %uint_1 %uint_1
+ %true = OpConstantTrue %bool
+ %3 = OpTypeFunction %void
+ %main = OpFunction %void None %3
+ %6 = OpLabel
+ %7 = )" + test_case.opcode +
+ " %" + test_case.resultType + " %scope_subgroup " +
+ operation + " %" + test_case.operand + R"( %uint_1
+ OpReturn
+ OpFunctionEnd;
+ )";
+ const auto result = SinglePassRunAndMatch<TrimCapabilitiesPass>(
+ kTest, /* do_validation= */ true);
+ EXPECT_EQ(std::get<1>(result), Pass::Status::SuccessWithChange);
+}
+
+INSTANTIATE_TEST_SUITE_P(
+ TrimCapabilitiesPassTestSubgroupClustered_Unsigned_I,
+ TrimCapabilitiesPassTestSubgroupClustered_Unsigned,
+ ::testing::Combine(::testing::ValuesIn(kSubgroupTestCases),
+ ::testing::Values("ClusteredReduce")),
+ [](const ::testing::TestParamInfo<
+ TrimCapabilitiesPassTestSubgroupClustered_Unsigned::ParamType>& info) {
+ return std::get<0>(info.param).opcode + "_" + std::get<1>(info.param);
+ });
+
} // namespace
} // namespace opt
} // namespace spvtools