@@ -39,8 +39,10 @@ std::string GenerateBFloatCode(const std::string& main_body) {
3939OpCapability Shader
4040OpCapability BFloat16TypeKHR
4141OpCapability AtomicFloat16AddEXT
42+ OpCapability FMAKHR
4243OpCapability GroupNonUniformShuffle
4344OpExtension "SPV_EXT_shader_atomic_float16_add"
45+ OpExtension "SPV_KHR_fma"
4446OpExtension "SPV_KHR_bfloat16"
4547%1 = OpExtInstImport "GLSL.std.450"
4648OpMemoryModel Logical GLSL450
@@ -73,6 +75,43 @@ OpFunctionEnd)";
7375 return prefix + main_body + suffix;
7476}
7577
78+ std::string GenerateBFloatCoopMatCode (const std::string& main_body) {
79+ const std::string prefix =
80+ R"(
81+ OpCapability Shader
82+ OpCapability VulkanMemoryModel
83+ OpCapability BFloat16TypeKHR
84+ OpCapability BFloat16CooperativeMatrixKHR
85+ OpCapability CooperativeMatrixKHR
86+ OpExtension "SPV_KHR_bfloat16"
87+ OpExtension "SPV_KHR_cooperative_matrix"
88+ OpExtension "SPV_KHR_vulkan_memory_model"
89+ OpMemoryModel Logical Vulkan
90+ OpEntryPoint GLCompute %main "main"
91+ OpExecutionMode %main LocalSize 32 1 1
92+ OpSource GLSL 450
93+ OpName %main "main"
94+ %void = OpTypeVoid
95+ %bfloat16 = OpTypeFloat 16 BFloat16KHR
96+ %func = OpTypeFunction %void
97+ %u32 = OpTypeInt 32 0
98+ %u32_8 = OpConstant %u32 8
99+ %subgroup = OpConstant %u32 3
100+ %useA = OpConstant %u32 0
101+ %bf16_1 = OpConstant %bfloat16 1
102+ %bf16matA = OpTypeCooperativeMatrixKHR %bfloat16 %subgroup %u32_8 %u32_8 %useA
103+ %bf16mat_A_1 = OpConstantComposite %bf16matA %bf16_1
104+ %main = OpFunction %void None %func
105+ %main_entry = OpLabel)" ;
106+
107+ const std::string suffix =
108+ R"(
109+ OpReturn
110+ OpFunctionEnd)" ;
111+
112+ return prefix + main_body + suffix;
113+ }
114+
76115TEST_F (ValidateInvalidType, Bfloat16InvalidArithmeticInstruction) {
77116 const std::string body = R"(
78117%v1 = OpVariable %_ptr_Function_bfloat16 Function
@@ -89,6 +128,45 @@ TEST_F(ValidateInvalidType, Bfloat16InvalidArithmeticInstruction) {
89128 HasSubstr (" FMul doesn't support BFloat16 type." ));
90129}
91130
131+ TEST_F (ValidateInvalidType, Bfloat16InvalidFmaInstruction) {
132+ const std::string body = R"(
133+ %15 = OpFmaKHR %bfloat16 %bf16_1 %bf16_1 %bf16_1
134+ )" ;
135+
136+ CompileSuccessfully (GenerateBFloatCode (body).c_str (), SPV_ENV_VULKAN_1_3 );
137+ ASSERT_EQ (SPV_ERROR_INVALID_DATA ,
138+ ValidateInstructions (SPV_ENV_UNIVERSAL_1_6 ));
139+ EXPECT_THAT (getDiagnosticString (),
140+ HasSubstr (" FmaKHR doesn't support BFloat16 type." ));
141+ }
142+
143+ TEST_F (ValidateInvalidType, Bfloat16InvalidVectorTimesScalarInstruction) {
144+ const std::string body = R"(
145+ %v1 = OpVariable %_ptr_Function_v2bfloat16 Function
146+ %12 = OpLoad %v2bfloat16 %v1
147+ %15 = OpVectorTimesScalar %v2bfloat16 %12 %bf16_1
148+ )" ;
149+
150+ CompileSuccessfully (GenerateBFloatCode (body).c_str (), SPV_ENV_VULKAN_1_3 );
151+ ASSERT_EQ (SPV_ERROR_INVALID_DATA ,
152+ ValidateInstructions (SPV_ENV_UNIVERSAL_1_6 ));
153+ EXPECT_THAT (getDiagnosticString (),
154+ HasSubstr (" VectorTimesScalar doesn't support BFloat16 type." ));
155+ }
156+
157+ TEST_F (ValidateInvalidType, Bfloat16InvalidMatrixTimesScalarInstruction) {
158+ const std::string body = R"(
159+ %15 = OpMatrixTimesScalar %bf16matA %bf16mat_A_1 %bf16_1
160+ )" ;
161+
162+ CompileSuccessfully (GenerateBFloatCoopMatCode (body).c_str (),
163+ SPV_ENV_UNIVERSAL_1_6 );
164+ ASSERT_EQ (SPV_ERROR_INVALID_DATA ,
165+ ValidateInstructions (SPV_ENV_UNIVERSAL_1_6 ));
166+ EXPECT_THAT (getDiagnosticString (),
167+ HasSubstr (" MatrixTimesScalar doesn't support BFloat16 type." ));
168+ }
169+
92170TEST_F (ValidateInvalidType, Bfloat16InvalidRelationalInstruction) {
93171 const std::string body = R"(
94172%v1 = OpVariable %_ptr_Function_bfloat16 Function
@@ -209,6 +287,38 @@ TEST_F(ValidateInvalidType, FP8E5M2InvalidArithmeticInstruction) {
209287 HasSubstr (" FMul doesn't support FP8 E4M3/E5M2 types." ));
210288}
211289
290+ TEST_F (ValidateInvalidType, FP8E4M3InvalidDotInstruction) {
291+ const std::string body = R"(
292+ %v1 = OpVariable %_ptr_Function_v2fp8e4m3 Function
293+ %v2 = OpVariable %_ptr_Function_v2fp8e4m3 Function
294+ %12 = OpLoad %v2fp8e4m3 %v1
295+ %14 = OpLoad %v2fp8e4m3 %v2
296+ %15 = OpDot %fp8e4m3 %12 %14
297+ )" ;
298+
299+ CompileSuccessfully (GenerateFP8Code (body).c_str (), SPV_ENV_VULKAN_1_3 );
300+ ASSERT_EQ (SPV_ERROR_INVALID_DATA ,
301+ ValidateInstructions (SPV_ENV_UNIVERSAL_1_6 ));
302+ EXPECT_THAT (getDiagnosticString (),
303+ HasSubstr (" Dot doesn't support FP8 E4M3/E5M2 types." ));
304+ }
305+
306+ TEST_F (ValidateInvalidType, FP8E5M2InvalidDotInstruction) {
307+ const std::string body = R"(
308+ %v1 = OpVariable %_ptr_Function_v2fp8e5m2 Function
309+ %v2 = OpVariable %_ptr_Function_v2fp8e5m2 Function
310+ %12 = OpLoad %v2fp8e5m2 %v1
311+ %14 = OpLoad %v2fp8e5m2 %v2
312+ %15 = OpDot %fp8e5m2 %12 %14
313+ )" ;
314+
315+ CompileSuccessfully (GenerateFP8Code (body).c_str (), SPV_ENV_VULKAN_1_3 );
316+ ASSERT_EQ (SPV_ERROR_INVALID_DATA ,
317+ ValidateInstructions (SPV_ENV_UNIVERSAL_1_6 ));
318+ EXPECT_THAT (getDiagnosticString (),
319+ HasSubstr (" Dot doesn't support FP8 E4M3/E5M2 types." ));
320+ }
321+
212322TEST_F (ValidateInvalidType, FP8E4M3InvalidRelationalInstruction) {
213323 const std::string body = R"(
214324%v1 = OpVariable %_ptr_Function_fp8e4m3 Function
0 commit comments