/src/spirv-tools/source/val/validate_composites.cpp
Line | Count | Source |
1 | | // Copyright (c) 2017 Google Inc. |
2 | | // Modifications Copyright (C) 2024 Advanced Micro Devices, Inc. All rights |
3 | | // reserved. |
4 | | // |
5 | | // Licensed under the Apache License, Version 2.0 (the "License"); |
6 | | // you may not use this file except in compliance with the License. |
7 | | // You may obtain a copy of the License at |
8 | | // |
9 | | // http://www.apache.org/licenses/LICENSE-2.0 |
10 | | // |
11 | | // Unless required by applicable law or agreed to in writing, software |
12 | | // distributed under the License is distributed on an "AS IS" BASIS, |
13 | | // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. |
14 | | // See the License for the specific language governing permissions and |
15 | | // limitations under the License. |
16 | | |
17 | | // Validates correctness of composite SPIR-V instructions. |
18 | | |
19 | | #include <climits> |
20 | | #include <cstdint> |
21 | | |
22 | | #include "source/opcode.h" |
23 | | #include "source/spirv_target_env.h" |
24 | | #include "source/val/instruction.h" |
25 | | #include "source/val/validate.h" |
26 | | #include "source/val/validation_state.h" |
27 | | |
28 | | namespace spvtools { |
29 | | namespace val { |
30 | | namespace { |
31 | | |
32 | | // Returns the type of the value accessed by OpCompositeExtract or |
33 | | // OpCompositeInsert instruction. The function traverses the hierarchy of |
34 | | // nested data structures (structs, arrays, vectors, matrices) as directed by |
35 | | // the sequence of indices in the instruction. May return error if traversal |
36 | | // fails (encountered non-composite, out of bounds, no indices, nesting too |
37 | | // deep). |
38 | | spv_result_t GetExtractInsertValueType(ValidationState_t& _, |
39 | | const Instruction* inst, |
40 | | uint32_t* member_type, |
41 | 61.6k | uint32_t composite_id_index) { |
42 | 61.6k | const uint32_t num_operands = static_cast<uint32_t>(inst->operands().size()); |
43 | 61.6k | const uint32_t first_literal_index = composite_id_index + 1; |
44 | 61.6k | const uint32_t num_indices = num_operands - first_literal_index; |
45 | 61.6k | const uint32_t kCompositeExtractInsertMaxNumIndices = 255; |
46 | | |
47 | 61.6k | if (num_indices == 0) { |
48 | 11 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
49 | 11 | << "Expected at least one index to Op" |
50 | 11 | << spvOpcodeString(inst->opcode()) << ", zero found"; |
51 | | |
52 | 61.5k | } else if (num_indices > kCompositeExtractInsertMaxNumIndices) { |
53 | 1 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
54 | 1 | << "The number of indexes in Op" << spvOpcodeString(inst->opcode()) |
55 | 1 | << " may not exceed " << kCompositeExtractInsertMaxNumIndices |
56 | 1 | << ". Found " << num_indices << " indexes."; |
57 | 1 | } |
58 | | |
59 | 61.5k | *member_type = _.GetOperandTypeId(inst, composite_id_index); |
60 | 61.5k | if (*member_type == 0) { |
61 | 4 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
62 | 4 | << "Expected Composite to be an object of composite type"; |
63 | 4 | } |
64 | | |
65 | 61.5k | for (uint32_t operand_index = first_literal_index; |
66 | 123k | operand_index < num_operands; ++operand_index) { |
67 | 61.6k | const uint32_t component_index = |
68 | 61.6k | inst->GetOperandAs<uint32_t>(operand_index); |
69 | 61.6k | const Instruction* const type_inst = _.FindDef(*member_type); |
70 | 61.6k | assert(type_inst); |
71 | 61.6k | switch (type_inst->opcode()) { |
72 | 57.6k | case spv::Op::OpTypeVector: { |
73 | 57.6k | *member_type = type_inst->word(2); |
74 | 57.6k | const uint32_t vector_size = type_inst->word(3); |
75 | 57.6k | if (component_index >= vector_size) { |
76 | 60 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
77 | 60 | << "Vector access is out of bounds, vector size is " |
78 | 60 | << vector_size << ", but access index is " << component_index; |
79 | 60 | } |
80 | 57.5k | break; |
81 | 57.6k | } |
82 | 57.5k | case spv::Op::OpTypeMatrix: { |
83 | 257 | *member_type = type_inst->word(2); |
84 | 257 | const uint32_t num_cols = type_inst->word(3); |
85 | 257 | if (component_index >= num_cols) { |
86 | 15 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
87 | 15 | << "Matrix access is out of bounds, matrix has " << num_cols |
88 | 15 | << " columns, but access index is " << component_index; |
89 | 15 | } |
90 | 242 | break; |
91 | 257 | } |
92 | 242 | case spv::Op::OpTypeArray: { |
93 | 139 | uint64_t array_size = 0; |
94 | 139 | auto size = _.FindDef(type_inst->word(3)); |
95 | 139 | *member_type = type_inst->word(2); |
96 | 139 | if (spvOpcodeIsSpecConstant(size->opcode())) { |
97 | | // Cannot verify against the size of this array. |
98 | 61 | break; |
99 | 61 | } |
100 | | |
101 | 78 | if (!_.EvalConstantValUint64(type_inst->word(3), &array_size)) { |
102 | 0 | assert(0 && "Array type definition is corrupt"); |
103 | 0 | } |
104 | 78 | if (component_index >= array_size) { |
105 | 16 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
106 | 16 | << "Array access is out of bounds, array size is " |
107 | 16 | << array_size << ", but access index is " << component_index; |
108 | 16 | } |
109 | 62 | break; |
110 | 78 | } |
111 | 62 | case spv::Op::OpTypeRuntimeArray: |
112 | 4 | case spv::Op::OpTypeNodePayloadArrayAMDX: { |
113 | 4 | *member_type = type_inst->word(2); |
114 | | // Array size is unknown. |
115 | 4 | break; |
116 | 4 | } |
117 | 3.60k | case spv::Op::OpTypeStruct: { |
118 | 3.60k | const size_t num_struct_members = type_inst->words().size() - 2; |
119 | 3.60k | if (component_index >= num_struct_members) { |
120 | 10 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
121 | 10 | << "Index is out of bounds, can not find index " |
122 | 10 | << component_index << " in the structure <id> '" |
123 | 10 | << type_inst->id() << "'. This structure has " |
124 | 10 | << num_struct_members << " members. Largest valid index is " |
125 | 10 | << num_struct_members - 1 << "."; |
126 | 10 | } |
127 | 3.59k | *member_type = type_inst->word(component_index + 2); |
128 | 3.59k | break; |
129 | 3.60k | } |
130 | 0 | case spv::Op::OpTypeVectorIdEXT: |
131 | 0 | case spv::Op::OpTypeCooperativeMatrixKHR: |
132 | 0 | case spv::Op::OpTypeCooperativeMatrixNV: { |
133 | 0 | *member_type = type_inst->word(2); |
134 | 0 | break; |
135 | 0 | } |
136 | 17 | default: |
137 | 17 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
138 | 17 | << "Reached non-composite type while indexes still remain to " |
139 | 17 | "be traversed."; |
140 | 61.6k | } |
141 | 61.6k | } |
142 | | |
143 | 61.4k | return SPV_SUCCESS; |
144 | 61.5k | } |
145 | | |
146 | | spv_result_t ValidateVectorExtractDynamic(ValidationState_t& _, |
147 | 137 | const Instruction* inst) { |
148 | 137 | const uint32_t result_type = inst->type_id(); |
149 | 137 | const spv::Op result_opcode = _.GetIdOpcode(result_type); |
150 | 137 | if (!spvOpcodeIsScalarType(result_opcode)) { |
151 | 3 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
152 | 3 | << "Expected Result Type to be a scalar type"; |
153 | 3 | } |
154 | | |
155 | 134 | const uint32_t vector_type = _.GetOperandTypeId(inst, 2); |
156 | 134 | const spv::Op vector_opcode = _.GetIdOpcode(vector_type); |
157 | 134 | if (vector_opcode != spv::Op::OpTypeVector && |
158 | 4 | vector_opcode != spv::Op::OpTypeVectorIdEXT) { |
159 | 4 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
160 | 4 | << "Expected Vector type to be OpTypeVector"; |
161 | 4 | } |
162 | | |
163 | 130 | if (_.GetComponentType(vector_type) != result_type) { |
164 | 4 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
165 | 4 | << "Expected Vector component type to be equal to Result Type"; |
166 | 4 | } |
167 | | |
168 | 126 | const auto index = _.FindDef(inst->GetOperandAs<uint32_t>(3)); |
169 | 126 | if (!index || index->type_id() == 0 || !_.IsIntScalarType(index->type_id())) { |
170 | 8 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
171 | 8 | << "Expected Index to be int scalar"; |
172 | 8 | } |
173 | | |
174 | 118 | if (_.HasCapability(spv::Capability::Shader) && |
175 | 117 | _.ContainsLimitedUseIntOrFloatType(inst->type_id())) { |
176 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
177 | 0 | << "Cannot extract from a vector of 8- or 16-bit types"; |
178 | 0 | } |
179 | 118 | return SPV_SUCCESS; |
180 | 118 | } |
181 | | |
182 | | spv_result_t ValidateVectorInsertDyanmic(ValidationState_t& _, |
183 | 45 | const Instruction* inst) { |
184 | 45 | const uint32_t result_type = inst->type_id(); |
185 | 45 | const spv::Op result_opcode = _.GetIdOpcode(result_type); |
186 | 45 | if (result_opcode != spv::Op::OpTypeVector && |
187 | 4 | result_opcode != spv::Op::OpTypeVectorIdEXT) { |
188 | 4 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
189 | 4 | << "Expected Result Type to be OpTypeVector"; |
190 | 4 | } |
191 | | |
192 | 41 | const uint32_t vector_type = _.GetOperandTypeId(inst, 2); |
193 | 41 | if (vector_type != result_type) { |
194 | 4 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
195 | 4 | << "Expected Vector type to be equal to Result Type"; |
196 | 4 | } |
197 | | |
198 | 37 | const uint32_t component_type = _.GetOperandTypeId(inst, 3); |
199 | 37 | if (_.GetComponentType(result_type) != component_type) { |
200 | 3 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
201 | 3 | << "Expected Component type to be equal to Result Type " |
202 | 3 | << "component type"; |
203 | 3 | } |
204 | | |
205 | 34 | const uint32_t index_type = _.GetOperandTypeId(inst, 4); |
206 | 34 | if (!_.IsIntScalarType(index_type)) { |
207 | 3 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
208 | 3 | << "Expected Index to be int scalar"; |
209 | 3 | } |
210 | | |
211 | 31 | if (_.HasCapability(spv::Capability::Shader) && |
212 | 31 | _.ContainsLimitedUseIntOrFloatType(inst->type_id())) { |
213 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
214 | 0 | << "Cannot insert into a vector of 8- or 16-bit types"; |
215 | 0 | } |
216 | 31 | return SPV_SUCCESS; |
217 | 31 | } |
218 | | |
219 | | spv_result_t ValidateCompositeConstruct(ValidationState_t& _, |
220 | 34.1k | const Instruction* inst) { |
221 | 34.1k | const uint32_t num_operands = static_cast<uint32_t>(inst->operands().size()); |
222 | 34.1k | const uint32_t result_type = inst->type_id(); |
223 | 34.1k | const spv::Op result_opcode = _.GetIdOpcode(result_type); |
224 | 34.1k | switch (result_opcode) { |
225 | 32.2k | case spv::Op::OpTypeVector: |
226 | 32.2k | case spv::Op::OpTypeVectorIdEXT: { |
227 | 32.2k | uint32_t num_result_components = _.GetDimension(result_type); |
228 | 32.2k | const uint32_t result_component_type = _.GetComponentType(result_type); |
229 | 32.2k | uint32_t given_component_count = 0; |
230 | | |
231 | 32.2k | bool comp_is_int32 = true, comp_is_const_int32 = true; |
232 | | |
233 | 32.2k | if (result_opcode == spv::Op::OpTypeVector) { |
234 | 32.2k | if (num_operands <= 3 && |
235 | 3 | !_.HasCapability(spv::Capability::LongVectorEXT)) { |
236 | 3 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
237 | 3 | << "Expected number of constituents to be at least 2"; |
238 | 3 | } |
239 | 32.2k | } else { |
240 | 0 | uint32_t comp_count_id = |
241 | 0 | _.FindDef(result_type)->GetOperandAs<uint32_t>(2); |
242 | 0 | std::tie(comp_is_int32, comp_is_const_int32, num_result_components) = |
243 | 0 | _.EvalInt32IfConst(comp_count_id); |
244 | 0 | } |
245 | | |
246 | 126k | for (uint32_t operand_index = 2; operand_index < num_operands; |
247 | 94.1k | ++operand_index) { |
248 | 94.1k | const uint32_t operand_type = _.GetOperandTypeId(inst, operand_index); |
249 | 94.1k | if (operand_type == result_component_type) { |
250 | 94.0k | ++given_component_count; |
251 | 94.0k | } else { |
252 | 96 | if (_.GetIdOpcode(operand_type) != spv::Op::OpTypeVector || |
253 | 72 | _.GetComponentType(operand_type) != result_component_type) { |
254 | 28 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
255 | 28 | << "Expected Constituents to be scalars or vectors of" |
256 | 28 | << " the same type as Result Type components"; |
257 | 28 | } |
258 | | |
259 | 68 | given_component_count += _.GetDimension(operand_type); |
260 | 68 | } |
261 | 94.1k | } |
262 | | |
263 | 32.2k | if (comp_is_const_int32 && |
264 | 32.2k | num_result_components != given_component_count) { |
265 | 13 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
266 | 13 | << "Expected total number of given components to be equal " |
267 | 13 | << "to the size of Result Type vector"; |
268 | 13 | } |
269 | | |
270 | 32.2k | break; |
271 | 32.2k | } |
272 | 32.2k | case spv::Op::OpTypeMatrix: { |
273 | 839 | uint32_t result_num_rows = 0; |
274 | 839 | uint32_t result_num_cols = 0; |
275 | 839 | uint32_t result_col_type = 0; |
276 | 839 | uint32_t result_component_type = 0; |
277 | 839 | if (!_.GetMatrixTypeInfo(result_type, &result_num_rows, &result_num_cols, |
278 | 839 | &result_col_type, &result_component_type)) { |
279 | 0 | assert(0); |
280 | 0 | } |
281 | | |
282 | 839 | if (result_num_cols + 2 != num_operands) { |
283 | 13 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
284 | 13 | << "Expected total number of Constituents to be equal " |
285 | 13 | << "to the number of columns of Result Type matrix"; |
286 | 13 | } |
287 | | |
288 | 2.67k | for (uint32_t operand_index = 2; operand_index < num_operands; |
289 | 1.86k | ++operand_index) { |
290 | 1.86k | const uint32_t operand_type = _.GetOperandTypeId(inst, operand_index); |
291 | 1.86k | if (operand_type != result_col_type) { |
292 | 11 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
293 | 11 | << "Expected Constituent type to be equal to the column " |
294 | 11 | << "type Result Type matrix"; |
295 | 11 | } |
296 | 1.86k | } |
297 | | |
298 | 815 | break; |
299 | 826 | } |
300 | 815 | case spv::Op::OpTypeArray: { |
301 | 73 | const Instruction* const array_inst = _.FindDef(result_type); |
302 | 73 | assert(array_inst); |
303 | 73 | assert(array_inst->opcode() == spv::Op::OpTypeArray); |
304 | | |
305 | 73 | auto size = _.FindDef(array_inst->word(3)); |
306 | 73 | if (spvOpcodeIsSpecConstant(size->opcode())) { |
307 | | // Cannot verify against the size of this array. |
308 | 5 | break; |
309 | 5 | } |
310 | | |
311 | 68 | uint64_t array_size = 0; |
312 | 68 | if (!_.EvalConstantValUint64(array_inst->word(3), &array_size)) { |
313 | 0 | assert(0 && "Array type definition is corrupt"); |
314 | 0 | } |
315 | | |
316 | 68 | if (array_size + 2 != num_operands) { |
317 | 16 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
318 | 16 | << "Expected total number of Constituents to be equal " |
319 | 16 | << "to the number of elements of Result Type array"; |
320 | 16 | } |
321 | | |
322 | 52 | const uint32_t result_component_type = array_inst->word(2); |
323 | 849 | for (uint32_t operand_index = 2; operand_index < num_operands; |
324 | 822 | ++operand_index) { |
325 | 822 | const uint32_t operand_type = _.GetOperandTypeId(inst, operand_index); |
326 | 822 | if (operand_type != result_component_type) { |
327 | 25 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
328 | 25 | << "Expected Constituent type to be equal to the column " |
329 | 25 | << "type Result Type array"; |
330 | 25 | } |
331 | 822 | } |
332 | | |
333 | 27 | break; |
334 | 52 | } |
335 | 907 | case spv::Op::OpTypeStruct: { |
336 | 907 | const Instruction* const struct_inst = _.FindDef(result_type); |
337 | 907 | assert(struct_inst); |
338 | 907 | assert(struct_inst->opcode() == spv::Op::OpTypeStruct); |
339 | | |
340 | 907 | if (struct_inst->operands().size() + 1 != num_operands) { |
341 | 3 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
342 | 3 | << "Expected total number of Constituents to be equal " |
343 | 3 | << "to the number of members of Result Type struct"; |
344 | 3 | } |
345 | | |
346 | 1.86k | for (uint32_t operand_index = 2; operand_index < num_operands; |
347 | 971 | ++operand_index) { |
348 | 971 | const uint32_t operand_type = _.GetOperandTypeId(inst, operand_index); |
349 | 971 | const uint32_t member_type = struct_inst->word(operand_index); |
350 | 971 | if (operand_type != member_type) { |
351 | 7 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
352 | 7 | << "Expected Constituent type to be equal to the " |
353 | 7 | << "corresponding member type of Result Type struct"; |
354 | 7 | } |
355 | 971 | } |
356 | | |
357 | 897 | break; |
358 | 904 | } |
359 | 897 | case spv::Op::OpTypeCooperativeMatrixKHR: { |
360 | 0 | const auto result_type_inst = _.FindDef(result_type); |
361 | 0 | assert(result_type_inst); |
362 | 0 | const auto component_type_id = |
363 | 0 | result_type_inst->GetOperandAs<uint32_t>(1); |
364 | |
|
365 | 0 | if (3 != num_operands) { |
366 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
367 | 0 | << "Must be only one constituent"; |
368 | 0 | } |
369 | | |
370 | 0 | const uint32_t operand_type_id = _.GetOperandTypeId(inst, 2); |
371 | |
|
372 | 0 | if (operand_type_id != component_type_id) { |
373 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
374 | 0 | << "Expected Constituent type to be equal to the component type"; |
375 | 0 | } |
376 | 0 | break; |
377 | 0 | } |
378 | 0 | case spv::Op::OpTypeCooperativeMatrixNV: { |
379 | 0 | const auto result_type_inst = _.FindDef(result_type); |
380 | 0 | assert(result_type_inst); |
381 | 0 | const auto component_type_id = |
382 | 0 | result_type_inst->GetOperandAs<uint32_t>(1); |
383 | |
|
384 | 0 | if (3 != num_operands) { |
385 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
386 | 0 | << "Expected single constituent"; |
387 | 0 | } |
388 | | |
389 | 0 | const uint32_t operand_type_id = _.GetOperandTypeId(inst, 2); |
390 | |
|
391 | 0 | if (operand_type_id != component_type_id) { |
392 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
393 | 0 | << "Expected Constituent type to be equal to the component type"; |
394 | 0 | } |
395 | | |
396 | 0 | break; |
397 | 0 | } |
398 | 12 | default: { |
399 | 12 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
400 | 12 | << "Expected Result Type to be a composite type"; |
401 | 0 | } |
402 | 34.1k | } |
403 | | |
404 | 33.9k | if (_.HasCapability(spv::Capability::Shader) && |
405 | 33.9k | _.ContainsLimitedUseIntOrFloatType(inst->type_id())) { |
406 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
407 | 0 | << "Cannot create a composite containing 8- or 16-bit types"; |
408 | 0 | } |
409 | 33.9k | return SPV_SUCCESS; |
410 | 33.9k | } |
411 | | |
412 | | spv_result_t ValidateCompositeConstructReplicate(ValidationState_t& _, |
413 | 0 | const Instruction* inst) { |
414 | 0 | const auto result_type = _.FindDef(inst->type_id()); |
415 | 0 | const uint32_t operand_type = _.GetOperandTypeId(inst, 2); |
416 | |
|
417 | 0 | switch (result_type->opcode()) { |
418 | 0 | case spv::Op::OpTypeVector: |
419 | 0 | case spv::Op::OpTypeVectorIdEXT: |
420 | 0 | case spv::Op::OpTypeMatrix: |
421 | 0 | case spv::Op::OpTypeArray: |
422 | 0 | case spv::Op::OpTypeCooperativeMatrixKHR: |
423 | 0 | case spv::Op::OpTypeCooperativeMatrixNV: { |
424 | 0 | const auto element_type = result_type->GetOperandAs<uint32_t>(1); |
425 | 0 | if (operand_type != element_type) { |
426 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
427 | 0 | << "Expected Value type to be equal to the " |
428 | 0 | << "result's element type"; |
429 | 0 | } |
430 | 0 | break; |
431 | 0 | } |
432 | 0 | case spv::Op::OpTypeStruct: { |
433 | 0 | for (uint32_t operand_index = 1; |
434 | 0 | operand_index < result_type->operands().size(); ++operand_index) { |
435 | 0 | const uint32_t member_type = |
436 | 0 | result_type->GetOperandAs<uint32_t>(operand_index); |
437 | 0 | if (operand_type != member_type) { |
438 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
439 | 0 | << "Expected Value type to be equal to the " |
440 | 0 | << "corresponding member type of the result"; |
441 | 0 | } |
442 | 0 | } |
443 | 0 | break; |
444 | 0 | } |
445 | 0 | case spv::Op::OpTypeTensorARM: { |
446 | 0 | const uint32_t component_type = result_type->GetOperandAs<uint32_t>(1); |
447 | 0 | if (operand_type != component_type) { |
448 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
449 | 0 | << "Expected Value type to be equal to the result's element " |
450 | 0 | "type"; |
451 | 0 | } |
452 | 0 | if (result_type->operands().size() <= 3) { |
453 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
454 | 0 | << "Result tensor type is not a composite type because it lacks " |
455 | 0 | "a shape operand"; |
456 | 0 | } |
457 | 0 | break; |
458 | 0 | } |
459 | 0 | default: { |
460 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
461 | 0 | << "Expected Result Type to be a composite type"; |
462 | 0 | } |
463 | 0 | } |
464 | | |
465 | 0 | if (_.HasCapability(spv::Capability::Shader) && |
466 | 0 | _.ContainsLimitedUseIntOrFloatType(inst->type_id())) { |
467 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
468 | 0 | << "Cannot create a composite containing 8- or 16-bit types"; |
469 | 0 | } |
470 | 0 | return SPV_SUCCESS; |
471 | 0 | } |
472 | | |
473 | | spv_result_t ValidateCompositeExtract(ValidationState_t& _, |
474 | | const Instruction* inst, |
475 | 50.9k | uint32_t operand_index = 2) { |
476 | 50.9k | uint32_t member_type = 0; |
477 | | |
478 | 50.9k | if (spv_result_t error = |
479 | 50.9k | GetExtractInsertValueType(_, inst, &member_type, operand_index)) { |
480 | 124 | return error; |
481 | 124 | } |
482 | | |
483 | 50.7k | const uint32_t result_type = inst->type_id(); |
484 | 50.7k | if (result_type != member_type) { |
485 | 6 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
486 | 6 | << "Result type (Op" << spvOpcodeString(_.GetIdOpcode(result_type)) |
487 | 6 | << ") does not match the type that results from indexing into " |
488 | 6 | "the composite (Op" |
489 | 6 | << spvOpcodeString(_.GetIdOpcode(member_type)) << ")."; |
490 | 6 | } |
491 | | |
492 | 50.7k | if (_.HasCapability(spv::Capability::Shader) && |
493 | 50.7k | _.ContainsLimitedUseIntOrFloatType(inst->type_id())) { |
494 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
495 | 0 | << "Cannot extract from a composite of 8- or 16-bit types"; |
496 | 0 | } |
497 | | |
498 | 50.7k | return SPV_SUCCESS; |
499 | 50.7k | } |
500 | | |
501 | | spv_result_t ValidateCompositeInsert(ValidationState_t& _, |
502 | | const Instruction* inst, |
503 | 10.7k | uint32_t operand_index = 2) { |
504 | 10.7k | const uint32_t object_type = _.GetOperandTypeId(inst, operand_index); |
505 | 10.7k | const uint32_t composite_type = _.GetOperandTypeId(inst, operand_index + 1); |
506 | 10.7k | const uint32_t result_type = inst->type_id(); |
507 | 10.7k | if (result_type != composite_type) { |
508 | 9 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
509 | 9 | << "The Result Type must be the same as Composite type in Op" |
510 | 9 | << spvOpcodeString(inst->opcode()) << " yielding Result Id " |
511 | 9 | << result_type << "."; |
512 | 9 | } |
513 | | |
514 | 10.6k | uint32_t member_type = 0; |
515 | 10.6k | if (spv_result_t error = |
516 | 10.6k | GetExtractInsertValueType(_, inst, &member_type, operand_index + 1)) { |
517 | 10 | return error; |
518 | 10 | } |
519 | | |
520 | 10.6k | if (object_type != member_type) { |
521 | 4 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
522 | 4 | << "The Object type (Op" |
523 | 4 | << spvOpcodeString(_.GetIdOpcode(object_type)) |
524 | 4 | << ") does not match the type that results from indexing into the " |
525 | 4 | "Composite (Op" |
526 | 4 | << spvOpcodeString(_.GetIdOpcode(member_type)) << ")."; |
527 | 4 | } |
528 | | |
529 | 10.6k | if (_.HasCapability(spv::Capability::Shader) && |
530 | 10.6k | _.ContainsLimitedUseIntOrFloatType(inst->type_id())) { |
531 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
532 | 0 | << "Cannot insert into a composite of 8- or 16-bit types"; |
533 | 0 | } |
534 | | |
535 | 10.6k | return SPV_SUCCESS; |
536 | 10.6k | } |
537 | | |
538 | 1.17k | spv_result_t ValidateCopyObject(ValidationState_t& _, const Instruction* inst) { |
539 | 1.17k | const uint32_t result_type = inst->type_id(); |
540 | 1.17k | const uint32_t operand_type = _.GetOperandTypeId(inst, 2); |
541 | 1.17k | if (operand_type != result_type) { |
542 | 8 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
543 | 8 | << "Expected Result Type and Operand type to be the same"; |
544 | 8 | } |
545 | 1.16k | if (_.IsVoidType(result_type)) { |
546 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
547 | 0 | << "OpCopyObject cannot have void result type"; |
548 | 0 | } |
549 | 1.16k | return SPV_SUCCESS; |
550 | 1.16k | } |
551 | | |
552 | 63 | spv_result_t ValidateTranspose(ValidationState_t& _, const Instruction* inst) { |
553 | 63 | uint32_t result_num_rows = 0; |
554 | 63 | uint32_t result_num_cols = 0; |
555 | 63 | uint32_t result_col_type = 0; |
556 | 63 | uint32_t result_component_type = 0; |
557 | 63 | const uint32_t result_type = inst->type_id(); |
558 | 63 | if (!_.GetMatrixTypeInfo(result_type, &result_num_rows, &result_num_cols, |
559 | 63 | &result_col_type, &result_component_type)) { |
560 | 9 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
561 | 9 | << "Expected Result Type to be a matrix type"; |
562 | 9 | } |
563 | | |
564 | 54 | const uint32_t matrix_type = _.GetOperandTypeId(inst, 2); |
565 | 54 | uint32_t matrix_num_rows = 0; |
566 | 54 | uint32_t matrix_num_cols = 0; |
567 | 54 | uint32_t matrix_col_type = 0; |
568 | 54 | uint32_t matrix_component_type = 0; |
569 | 54 | if (!_.GetMatrixTypeInfo(matrix_type, &matrix_num_rows, &matrix_num_cols, |
570 | 54 | &matrix_col_type, &matrix_component_type)) { |
571 | 4 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
572 | 4 | << "Expected Matrix to be of type OpTypeMatrix"; |
573 | 4 | } |
574 | | |
575 | 50 | if (result_component_type != matrix_component_type) { |
576 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
577 | 0 | << "Expected component types of Matrix and Result Type to be " |
578 | 0 | << "identical"; |
579 | 0 | } |
580 | | |
581 | 50 | if (result_num_rows != matrix_num_cols || |
582 | 46 | result_num_cols != matrix_num_rows) { |
583 | 4 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
584 | 4 | << "Expected number of columns and the column size of Matrix " |
585 | 4 | << "to be the reverse of those of Result Type"; |
586 | 4 | } |
587 | | |
588 | 46 | if (_.HasCapability(spv::Capability::Shader) && |
589 | 46 | _.ContainsLimitedUseIntOrFloatType(inst->type_id())) { |
590 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
591 | 0 | << "Cannot transpose matrices of 16-bit floats"; |
592 | 0 | } |
593 | 46 | return SPV_SUCCESS; |
594 | 46 | } |
595 | | |
596 | | spv_result_t ValidateVectorShuffle(ValidationState_t& _, |
597 | | const Instruction* inst, |
598 | 14.8k | uint32_t operand_index = 2) { |
599 | 14.8k | auto result_type = _.FindDef(inst->type_id()); |
600 | 14.8k | if (!_.IsVectorType(result_type->id())) { |
601 | 11 | return _.diag(SPV_ERROR_INVALID_ID, inst) |
602 | 11 | << "The Result Type of OpVectorShuffle must be" |
603 | 11 | << " a vector type. Found Op" |
604 | 11 | << spvOpcodeString(result_type->opcode()) << "."; |
605 | 11 | } |
606 | | |
607 | | // The number of components in Result Type must be the same as the number of |
608 | | // Component operands. |
609 | 14.8k | uint32_t first_literal_index = operand_index + 2; |
610 | 14.8k | uint32_t component_count = |
611 | 14.8k | static_cast<uint32_t>(inst->operands().size()) - first_literal_index; |
612 | 14.8k | auto result_vec_dimension = _.GetDimension(result_type->id()); |
613 | 14.8k | if (result_vec_dimension > 0 && component_count != result_vec_dimension) { |
614 | 12 | return _.diag(SPV_ERROR_INVALID_ID, inst) |
615 | 12 | << "OpVectorShuffle component literals count does not match " |
616 | 12 | "Result Type <id> " |
617 | 12 | << _.getIdName(result_type->id()) << "s vector component count."; |
618 | 12 | } |
619 | | |
620 | | // Vector 1 and Vector 2 must both have vector types, with the same Component |
621 | | // Type as Result Type. |
622 | 14.8k | auto vec1_type = _.FindDef(_.GetOperandTypeId(inst, operand_index)); |
623 | 14.8k | auto vec2_type = _.FindDef(_.GetOperandTypeId(inst, operand_index + 1)); |
624 | 14.8k | if (!vec1_type || !_.IsVectorType(vec1_type->id())) { |
625 | 11 | return _.diag(SPV_ERROR_INVALID_ID, inst) |
626 | 11 | << "The type of Vector 1 must be a vector type."; |
627 | 11 | } |
628 | 14.8k | if (!vec2_type || !_.IsVectorType(vec2_type->id())) { |
629 | 9 | return _.diag(SPV_ERROR_INVALID_ID, inst) |
630 | 9 | << "The type of Vector 2 must be a vector type."; |
631 | 9 | } |
632 | | |
633 | 14.8k | uint32_t result_component_type = result_type->GetOperandAs<uint32_t>(1); |
634 | 14.8k | if (vec1_type->GetOperandAs<uint32_t>(1) != result_component_type) { |
635 | 3 | return _.diag(SPV_ERROR_INVALID_ID, inst) |
636 | 3 | << "The Component Type of Vector 1 must be the same as ResultType."; |
637 | 3 | } |
638 | 14.8k | if (vec2_type->GetOperandAs<uint32_t>(1) != result_component_type) { |
639 | 4 | return _.diag(SPV_ERROR_INVALID_ID, inst) |
640 | 4 | << "The Component Type of Vector 2 must be the same as ResultType."; |
641 | 4 | } |
642 | | |
643 | | // All Component literals must either be FFFFFFFF or in [0, N - 1]. |
644 | 14.8k | uint32_t vec1_component_count = vec1_type->GetOperandAs<uint32_t>(2); |
645 | 14.8k | uint32_t vec2_component_count = vec2_type->GetOperandAs<uint32_t>(2); |
646 | 14.8k | uint32_t N = vec1_component_count + vec2_component_count; |
647 | 47.8k | for (size_t i = first_literal_index; i < inst->operands().size(); ++i) { |
648 | 33.0k | uint32_t literal = inst->GetOperandAs<uint32_t>(i); |
649 | 33.0k | if (literal != 0xFFFFFFFF && literal >= N) { |
650 | 79 | return _.diag(SPV_ERROR_INVALID_ID, inst) |
651 | 79 | << "Component index " << literal << " is out of bounds for " |
652 | 79 | << "combined (Vector1 + Vector2) size of " << N << "."; |
653 | 79 | } |
654 | 33.0k | } |
655 | | |
656 | 14.7k | if (_.HasCapability(spv::Capability::Shader) && |
657 | 14.7k | _.ContainsLimitedUseIntOrFloatType(inst->type_id())) { |
658 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
659 | 0 | << "Cannot shuffle a vector of 8- or 16-bit types"; |
660 | 0 | } |
661 | | |
662 | 14.7k | return SPV_SUCCESS; |
663 | 14.7k | } |
664 | | |
665 | | spv_result_t ValidateCopyLogical(ValidationState_t& _, |
666 | 0 | const Instruction* inst) { |
667 | 0 | const auto result_type = _.FindDef(inst->type_id()); |
668 | 0 | const auto source = _.FindDef(inst->GetOperandAs<uint32_t>(2u)); |
669 | 0 | const auto source_type = _.FindDef(source->type_id()); |
670 | 0 | if (!source_type || !result_type || source_type == result_type) { |
671 | 0 | return _.diag(SPV_ERROR_INVALID_ID, inst) |
672 | 0 | << "Result Type must not equal the Operand type"; |
673 | 0 | } |
674 | | |
675 | 0 | if (!_.LogicallyMatch(source_type, result_type, false)) { |
676 | 0 | return _.diag(SPV_ERROR_INVALID_ID, inst) |
677 | 0 | << "Result Type does not logically match the Operand type"; |
678 | 0 | } |
679 | | |
680 | 0 | if (_.HasCapability(spv::Capability::Shader) && |
681 | 0 | _.ContainsLimitedUseIntOrFloatType(inst->type_id())) { |
682 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
683 | 0 | << "Cannot copy composites of 8- or 16-bit types"; |
684 | 0 | } |
685 | | |
686 | 0 | return SPV_SUCCESS; |
687 | 0 | } |
688 | | |
689 | | spv_result_t ValidateCompositeConstructCoopMatQCOM(ValidationState_t& _, |
690 | 0 | const Instruction* inst) { |
691 | | // Is the result of coop mat ? |
692 | 0 | const auto result_type_inst = _.FindDef(inst->type_id()); |
693 | 0 | if (!result_type_inst || |
694 | 0 | result_type_inst->opcode() != spv::Op::OpTypeCooperativeMatrixKHR) { |
695 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
696 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
697 | 0 | << " requires the result type be OpTypeCooperativeMatrixKHR"; |
698 | 0 | } |
699 | | |
700 | 0 | const auto source = _.FindDef(inst->GetOperandAs<uint32_t>(2u)); |
701 | 0 | const auto source_type_inst = _.FindDef(source->type_id()); |
702 | |
|
703 | 0 | if (!source_type_inst || source_type_inst->opcode() != spv::Op::OpTypeArray) { |
704 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
705 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
706 | 0 | << " requires the input operand be an OpTypeArray."; |
707 | 0 | } |
708 | | |
709 | | // Is the scope Subgrouop ? |
710 | 0 | { |
711 | 0 | unsigned scope = UINT_MAX; |
712 | 0 | unsigned scope_id = result_type_inst->GetOperandAs<unsigned>(2u); |
713 | 0 | bool status = _.GetConstantValueAs<unsigned>(scope_id, scope); |
714 | 0 | bool is_scope_spec_const = |
715 | 0 | spvOpcodeIsSpecConstant(_.FindDef(scope_id)->opcode()); |
716 | 0 | if (!is_scope_spec_const && |
717 | 0 | (!status || scope != static_cast<uint64_t>(spv::Scope::Subgroup))) { |
718 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
719 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
720 | 0 | << " requires the result type's scope be Subgroup."; |
721 | 0 | } |
722 | 0 | } |
723 | | |
724 | 0 | unsigned ar_len = UINT_MAX; |
725 | 0 | unsigned src_arr_len_id = source_type_inst->GetOperandAs<unsigned>(2u); |
726 | 0 | bool ar_len_status = _.GetConstantValueAs<unsigned>(src_arr_len_id, ar_len); |
727 | 0 | bool is_src_arr_len_spec_const = |
728 | 0 | spvOpcodeIsSpecConstant(_.FindDef(src_arr_len_id)->opcode()); |
729 | |
|
730 | 0 | const auto source_elt_type = _.GetComponentType(source_type_inst->id()); |
731 | 0 | const auto result_elt_type = result_type_inst->GetOperandAs<uint32_t>(1u); |
732 | |
|
733 | 0 | if ((source_elt_type != result_elt_type) && |
734 | 0 | !(_.ContainsSizedIntOrFloatType(source_elt_type, spv::Op::OpTypeInt, |
735 | 0 | 32) && |
736 | 0 | _.IsUnsignedIntScalarType(source_elt_type))) { |
737 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
738 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
739 | 0 | << " requires ether the input element type is equal to the result " |
740 | 0 | "element type or it is the unsigned 32-bit integer."; |
741 | 0 | } |
742 | | |
743 | 0 | unsigned res_row_id = result_type_inst->GetOperandAs<unsigned>(3u); |
744 | 0 | unsigned res_col_id = result_type_inst->GetOperandAs<unsigned>(4u); |
745 | 0 | unsigned res_use_id = result_type_inst->GetOperandAs<unsigned>(5u); |
746 | |
|
747 | 0 | unsigned cm_use = UINT_MAX; |
748 | 0 | bool cm_use_status = _.GetConstantValueAs<unsigned>(res_use_id, cm_use); |
749 | |
|
750 | 0 | switch (static_cast<spv::CooperativeMatrixUse>(cm_use)) { |
751 | 0 | case spv::CooperativeMatrixUse::MatrixAKHR: { |
752 | | // result coopmat component type check |
753 | 0 | if (!_.IsIntNOrFP32OrFP16<8>(result_elt_type)) { |
754 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
755 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
756 | 0 | << " requires the result element type is one of 8-bit OpTypeInt " |
757 | 0 | "signed/unsigned, 16- or 32-bit OpTypeFloat" |
758 | 0 | << " when result coopmat's use is MatrixAKHR"; |
759 | 0 | } |
760 | | |
761 | | // result coopmat column length check |
762 | 0 | unsigned n_cols = UINT_MAX; |
763 | 0 | bool status = _.GetConstantValueAs<unsigned>(res_col_id, n_cols); |
764 | 0 | bool is_res_col_spec_const = |
765 | 0 | spvOpcodeIsSpecConstant(_.FindDef(res_col_id)->opcode()); |
766 | 0 | if (!is_res_col_spec_const && |
767 | 0 | (!status || (!(_.ContainsSizedIntOrFloatType(result_elt_type, |
768 | 0 | spv::Op::OpTypeInt, 8) && |
769 | 0 | n_cols == 32) && |
770 | 0 | !(_.ContainsSizedIntOrFloatType( |
771 | 0 | result_elt_type, spv::Op::OpTypeFloat, 16) && |
772 | 0 | n_cols == 16) && |
773 | 0 | !(_.ContainsSizedIntOrFloatType( |
774 | 0 | result_elt_type, spv::Op::OpTypeFloat, 32) && |
775 | 0 | n_cols == 8)))) { |
776 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
777 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
778 | 0 | << " requires the columns of the result coopmat have the bit " |
779 | 0 | "length of 256" |
780 | 0 | << " when result coopmat's use is MatrixAKHR"; |
781 | 0 | } |
782 | | // source array length check |
783 | 0 | if (!is_src_arr_len_spec_const && |
784 | 0 | (!ar_len_status || |
785 | 0 | (!(_.ContainsSizedIntOrFloatType(source_elt_type, spv::Op::OpTypeInt, |
786 | 0 | 32) && |
787 | 0 | _.IsUnsignedIntScalarType(source_elt_type) && (ar_len == 8)) && |
788 | 0 | !(n_cols == ar_len)))) { |
789 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
790 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
791 | 0 | << " requires the source array length be 8 if its elt type is " |
792 | 0 | "32-bit unsigned OpTypeInt and be the result's number of " |
793 | 0 | "columns, otherwise" |
794 | 0 | << " when result coopmat's use is MatrixAKHR"; |
795 | 0 | } |
796 | 0 | break; |
797 | 0 | } |
798 | 0 | case spv::CooperativeMatrixUse::MatrixBKHR: { |
799 | | // result coopmat component type check |
800 | 0 | if (!_.IsIntNOrFP32OrFP16<8>(result_elt_type)) { |
801 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
802 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
803 | 0 | << " requires the result element type is one of 8-bit OpTypeInt " |
804 | 0 | "signed/unsigned, 16- or 32-bit OpTypeFloat" |
805 | 0 | << " when result coopmat's use is MatrixBKHR"; |
806 | 0 | } |
807 | | |
808 | | // result coopmat row length check |
809 | 0 | unsigned n_rows = UINT_MAX; |
810 | 0 | bool status = _.GetConstantValueAs<unsigned>(res_row_id, n_rows); |
811 | 0 | bool is_res_row_spec_const = |
812 | 0 | spvOpcodeIsSpecConstant(_.FindDef(res_row_id)->opcode()); |
813 | 0 | if (!is_res_row_spec_const && |
814 | 0 | (!status || (!(_.ContainsSizedIntOrFloatType(result_elt_type, |
815 | 0 | spv::Op::OpTypeInt, 8) && |
816 | 0 | n_rows == 32) && |
817 | 0 | !(_.ContainsSizedIntOrFloatType( |
818 | 0 | result_elt_type, spv::Op::OpTypeFloat, 16) && |
819 | 0 | n_rows == 16) && |
820 | 0 | !(_.ContainsSizedIntOrFloatType( |
821 | 0 | result_elt_type, spv::Op::OpTypeFloat, 32) && |
822 | 0 | n_rows == 8)))) { |
823 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
824 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
825 | 0 | << " requires the rows of the result operand have the bit " |
826 | 0 | "length of 256" |
827 | 0 | << " when result coopmat's use is MatrixBKHR"; |
828 | 0 | } |
829 | | // source array length check |
830 | 0 | if (!is_src_arr_len_spec_const && |
831 | 0 | (!ar_len_status || |
832 | 0 | (!(_.ContainsSizedIntOrFloatType(source_elt_type, spv::Op::OpTypeInt, |
833 | 0 | 32) && |
834 | 0 | _.IsUnsignedIntScalarType(source_elt_type) && (ar_len == 8)) && |
835 | 0 | !(n_rows == ar_len)))) { |
836 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
837 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
838 | 0 | << " requires the source array length be 8 if its elt type is " |
839 | 0 | "32-bit unsigned OpTypeInt and be the result's number of " |
840 | 0 | "rows, otherwise" |
841 | 0 | << " when result coopmat's use is MatrixBKHR"; |
842 | 0 | } |
843 | 0 | break; |
844 | 0 | } |
845 | 0 | case spv::CooperativeMatrixUse::MatrixAccumulatorKHR: { |
846 | | // result coopmat component type check |
847 | 0 | if (!_.IsIntNOrFP32OrFP16<32>(result_elt_type)) { |
848 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
849 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
850 | 0 | << " requires the result element type is one of 32-bit " |
851 | 0 | "OpTypeInt signed/unsigned, 16- or 32-bit OpTypeFloat" |
852 | 0 | << " when result coopmat's use is MatrixAccumulatorKHR"; |
853 | 0 | } |
854 | | |
855 | | // source array length check |
856 | 0 | unsigned n_cols = UINT_MAX; |
857 | 0 | bool status = _.GetConstantValueAs<unsigned>(res_col_id, n_cols); |
858 | 0 | bool is_res_col_spec_const = |
859 | 0 | spvOpcodeIsSpecConstant(_.FindDef(res_col_id)->opcode()); |
860 | 0 | if (!is_res_col_spec_const && !is_src_arr_len_spec_const && |
861 | 0 | (!status || !ar_len_status || |
862 | 0 | (!(_.ContainsSizedIntOrFloatType(source_elt_type, spv::Op::OpTypeInt, |
863 | 0 | 32) && |
864 | 0 | _.IsUnsignedIntScalarType(source_elt_type) && |
865 | 0 | (_.ContainsSizedIntOrFloatType(result_elt_type, |
866 | 0 | spv::Op::OpTypeFloat, 16) |
867 | 0 | ? (n_cols / 2 == ar_len) |
868 | 0 | : n_cols == ar_len)) && |
869 | 0 | (n_cols != ar_len)))) { |
870 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
871 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
872 | 0 | << " requires the source array length be a half of the number " |
873 | 0 | "of columns of the resulting cooerative matrix if the " |
874 | 0 | "matrix's componet type is 16-bit OpTypeFloat and be equal " |
875 | 0 | "to the number of columns, otherwise," |
876 | 0 | << " when result coopmat's use is MatrixAccumulatorKHR"; |
877 | 0 | } |
878 | 0 | break; |
879 | 0 | } |
880 | 0 | default: { |
881 | 0 | bool is_cm_use_spec_const = |
882 | 0 | spvOpcodeIsSpecConstant(_.FindDef(res_use_id)->opcode()); |
883 | 0 | if (!is_cm_use_spec_const || !cm_use_status) { |
884 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
885 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
886 | 0 | << " requires the the resulting cooerative matrix's use be " |
887 | 0 | << " one of MatrixAKHR (== 0), MatrixBKHR (== 1), and " |
888 | 0 | "MatrixAccumulatorKHR (== 2)"; |
889 | 0 | } |
890 | 0 | break; |
891 | 0 | } |
892 | 0 | } |
893 | | |
894 | 0 | return SPV_SUCCESS; |
895 | 0 | } |
896 | | |
897 | | spv_result_t ValidateCompositeExtractCoopMatQCOM(ValidationState_t& _, |
898 | 0 | const Instruction* inst) { |
899 | 0 | const auto result_type_inst = _.FindDef(inst->type_id()); |
900 | 0 | if (!result_type_inst || result_type_inst->opcode() != spv::Op::OpTypeArray) { |
901 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
902 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
903 | 0 | << " requires the input operand be an OpTypeArray."; |
904 | 0 | } |
905 | | |
906 | 0 | const auto source = _.FindDef(inst->GetOperandAs<uint32_t>(2u)); |
907 | 0 | const auto source_type_inst = _.FindDef(source->type_id()); |
908 | | |
909 | | // Is the source of coop mat ? |
910 | 0 | if (!source_type_inst || |
911 | 0 | source_type_inst->opcode() != spv::Op::OpTypeCooperativeMatrixKHR) { |
912 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
913 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
914 | 0 | << " requires the source type be OpTypeCooperativeMatrixKHR"; |
915 | 0 | } |
916 | | |
917 | | // Is the scope Subgrouop ? |
918 | 0 | { |
919 | 0 | unsigned scope = UINT_MAX; |
920 | 0 | unsigned scope_id = source_type_inst->GetOperandAs<unsigned>(2u); |
921 | 0 | bool status = _.GetConstantValueAs<unsigned>(scope_id, scope); |
922 | 0 | bool is_scope_spec_const = |
923 | 0 | spvOpcodeIsSpecConstant(_.FindDef(scope_id)->opcode()); |
924 | 0 | if (!is_scope_spec_const && |
925 | 0 | (!status || scope != static_cast<uint64_t>(spv::Scope::Subgroup))) { |
926 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
927 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
928 | 0 | << " requires the source type's scope be Subgroup."; |
929 | 0 | } |
930 | 0 | } |
931 | | |
932 | 0 | unsigned ar_len = UINT_MAX; |
933 | 0 | unsigned res_arr_len_id = result_type_inst->GetOperandAs<unsigned>(2u); |
934 | 0 | bool ar_len_status = _.GetConstantValueAs<unsigned>(res_arr_len_id, ar_len); |
935 | 0 | bool is_res_arr_len_spec_const = |
936 | 0 | spvOpcodeIsSpecConstant(_.FindDef(res_arr_len_id)->opcode()); |
937 | |
|
938 | 0 | const auto source_elt_type = _.GetComponentType(source_type_inst->id()); |
939 | 0 | const auto result_elt_type = result_type_inst->GetOperandAs<uint32_t>(1u); |
940 | |
|
941 | 0 | unsigned src_row_id = source_type_inst->GetOperandAs<unsigned>(3u); |
942 | 0 | unsigned src_col_id = source_type_inst->GetOperandAs<unsigned>(4u); |
943 | 0 | unsigned src_use_id = source_type_inst->GetOperandAs<unsigned>(5u); |
944 | |
|
945 | 0 | unsigned cm_use = UINT_MAX; |
946 | 0 | bool cm_use_status = _.GetConstantValueAs<unsigned>(src_use_id, cm_use); |
947 | |
|
948 | 0 | switch (static_cast<spv::CooperativeMatrixUse>(cm_use)) { |
949 | 0 | case spv::CooperativeMatrixUse::MatrixAKHR: { |
950 | | // source coopmat component type check |
951 | 0 | if (!_.IsIntNOrFP32OrFP16<8>(source_elt_type)) { |
952 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
953 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
954 | 0 | << " requires the source element type be one of 8-bit OpTypeInt " |
955 | 0 | "signed/unsigned, 16- or 32-bit OpTypeFloat" |
956 | 0 | << " when source coopmat's use is MatrixAKHR"; |
957 | 0 | } |
958 | | |
959 | | // source coopmat column length check |
960 | 0 | unsigned n_cols = UINT_MAX; |
961 | 0 | bool status = _.GetConstantValueAs<unsigned>(src_col_id, n_cols); |
962 | 0 | bool is_src_col_spec_const = |
963 | 0 | spvOpcodeIsSpecConstant(_.FindDef(src_col_id)->opcode()); |
964 | 0 | if (!is_src_col_spec_const && |
965 | 0 | (!status || (!(_.ContainsSizedIntOrFloatType(source_elt_type, |
966 | 0 | spv::Op::OpTypeInt, 8) && |
967 | 0 | n_cols == 32) && |
968 | 0 | !(_.ContainsSizedIntOrFloatType( |
969 | 0 | source_elt_type, spv::Op::OpTypeFloat, 16) && |
970 | 0 | n_cols == 16) && |
971 | 0 | !(_.ContainsSizedIntOrFloatType( |
972 | 0 | source_elt_type, spv::Op::OpTypeFloat, 32) && |
973 | 0 | n_cols == 8)))) { |
974 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
975 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
976 | 0 | << " requires the columns of the source coopmat have the bit " |
977 | 0 | "length of 256" |
978 | 0 | << " when source coopmat's use is MatrixAKHR"; |
979 | 0 | } |
980 | | // result type check |
981 | 0 | if (!is_res_arr_len_spec_const && |
982 | 0 | !(source_elt_type == result_elt_type && (n_cols == ar_len)) && |
983 | 0 | !(_.ContainsSizedIntOrFloatType(result_elt_type, spv::Op::OpTypeInt, |
984 | 0 | 32) && |
985 | 0 | _.IsUnsignedIntScalarType(result_elt_type) && (ar_len == 8))) { |
986 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
987 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
988 | 0 | << " requires either the result element type be the same as the " |
989 | 0 | "source cooperative matrix's component type" |
990 | 0 | << " and its length be the same as the number of columns of the " |
991 | 0 | "matrix or the result element type be" |
992 | 0 | << " unsigned 32-bit OpTypeInt and the length be 8" |
993 | 0 | << " when source coopmat's use is MatrixAKHR"; |
994 | 0 | } |
995 | 0 | break; |
996 | 0 | } |
997 | 0 | case spv::CooperativeMatrixUse::MatrixBKHR: { |
998 | | // source coopmat component type check |
999 | 0 | if (!_.IsIntNOrFP32OrFP16<8>(source_elt_type)) { |
1000 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
1001 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
1002 | 0 | << " requires the source element type be one of 8-bit OpTypeInt " |
1003 | 0 | "signed/unsigned, 16- or 32-bit OpTypeFloat" |
1004 | 0 | << " when source coopmat's use is MatrixBKHR"; |
1005 | 0 | } |
1006 | | |
1007 | | // source coopmat row length check |
1008 | 0 | unsigned n_rows = UINT_MAX; |
1009 | 0 | bool status = _.GetConstantValueAs<unsigned>(src_row_id, n_rows); |
1010 | 0 | bool is_src_row_spec_const = |
1011 | 0 | spvOpcodeIsSpecConstant(_.FindDef(src_row_id)->opcode()); |
1012 | 0 | if (!is_src_row_spec_const && |
1013 | 0 | (!status || (!(_.ContainsSizedIntOrFloatType(source_elt_type, |
1014 | 0 | spv::Op::OpTypeInt, 8) && |
1015 | 0 | n_rows == 32) && |
1016 | 0 | !(_.ContainsSizedIntOrFloatType( |
1017 | 0 | source_elt_type, spv::Op::OpTypeFloat, 16) && |
1018 | 0 | n_rows == 16) && |
1019 | 0 | !(_.ContainsSizedIntOrFloatType( |
1020 | 0 | source_elt_type, spv::Op::OpTypeFloat, 32) && |
1021 | 0 | n_rows == 8)))) { |
1022 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
1023 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
1024 | 0 | << " requires the rows of the source coopmat have the bit " |
1025 | 0 | "length of 256" |
1026 | 0 | << " when source coopmat's use is MatrixBKHR"; |
1027 | 0 | } |
1028 | | // result type check |
1029 | 0 | if (!is_res_arr_len_spec_const && |
1030 | 0 | !(source_elt_type == result_elt_type && (n_rows == ar_len)) && |
1031 | 0 | !(_.ContainsSizedIntOrFloatType(result_elt_type, spv::Op::OpTypeInt, |
1032 | 0 | 32) && |
1033 | 0 | _.IsUnsignedIntScalarType(result_elt_type) && (ar_len == 8))) { |
1034 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
1035 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
1036 | 0 | << " requires either the result element type be the same as the " |
1037 | 0 | "source cooperative matrix's component type" |
1038 | 0 | << " and its length be the same as the number of rows of the " |
1039 | 0 | "matrix or the result element type be" |
1040 | 0 | << " unsigned 32-bit OpTypeInt and the length be 8" |
1041 | 0 | << " when source coopmat's use is MatrixBKHR"; |
1042 | 0 | } |
1043 | 0 | break; |
1044 | 0 | } |
1045 | 0 | case spv::CooperativeMatrixUse::MatrixAccumulatorKHR: { |
1046 | | // source coopmat component type check |
1047 | 0 | if (!_.IsIntNOrFP32OrFP16<32>(source_elt_type)) { |
1048 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
1049 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
1050 | 0 | << " requires the source element type be one of 32-bit " |
1051 | 0 | "OpTypeInt signed/unsigned, 16- or 32-bit OpTypeFloat" |
1052 | 0 | << " when source coopmat's use is MatrixAccumulatorKHR"; |
1053 | 0 | } |
1054 | | |
1055 | | // result type check |
1056 | 0 | unsigned n_cols = UINT_MAX; |
1057 | 0 | bool status = _.GetConstantValueAs<unsigned>(src_col_id, n_cols); |
1058 | 0 | bool is_src_col_spec_const = |
1059 | 0 | spvOpcodeIsSpecConstant(_.FindDef(src_col_id)->opcode()); |
1060 | 0 | if (!is_src_col_spec_const && !is_res_arr_len_spec_const && |
1061 | 0 | (!status || !ar_len_status || |
1062 | 0 | (!(source_elt_type == result_elt_type && (n_cols == ar_len)) && |
1063 | 0 | !(_.ContainsSizedIntOrFloatType(result_elt_type, spv::Op::OpTypeInt, |
1064 | 0 | 32) && |
1065 | 0 | _.IsUnsignedIntScalarType(result_elt_type) && |
1066 | 0 | (_.ContainsSizedIntOrFloatType(source_elt_type, |
1067 | 0 | spv::Op::OpTypeFloat, 16) |
1068 | 0 | ? (n_cols / 2 == ar_len) |
1069 | 0 | : (n_cols == ar_len)))))) { |
1070 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
1071 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
1072 | 0 | << " requires either the result element type be the same as the " |
1073 | 0 | "source cooperative matrix's component type" |
1074 | 0 | << " and its length be the same as the number of columns of the " |
1075 | 0 | "matrix or the result element type be" |
1076 | 0 | << " unsigned 32-bit OpTypeInt and the length be the number of " |
1077 | 0 | "the columns of the matrix if its component" |
1078 | 0 | << " type is 32-bit OpTypeFloat and be a half of the number of " |
1079 | 0 | "the columns of the matrix if its component" |
1080 | 0 | << " type is 16-bit OpTypeFloat" |
1081 | 0 | << " when source coopmat's use is MatrixAccumulatorKHR"; |
1082 | 0 | } |
1083 | 0 | break; |
1084 | 0 | } |
1085 | 0 | default: { |
1086 | 0 | bool is_cm_use_spec_const = |
1087 | 0 | spvOpcodeIsSpecConstant(_.FindDef(src_use_id)->opcode()); |
1088 | 0 | if (!is_cm_use_spec_const || !cm_use_status) { |
1089 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
1090 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
1091 | 0 | << " requires the the source cooerative matrix's use be " |
1092 | 0 | << " one of MatrixAKHR (== 0), MatrixBKHR (== 1), and " |
1093 | 0 | "MatrixAccumulatorKHR (== 2)"; |
1094 | 0 | } |
1095 | 0 | break; |
1096 | 0 | } |
1097 | 0 | } |
1098 | | |
1099 | 0 | return SPV_SUCCESS; |
1100 | 0 | } |
1101 | | |
1102 | | spv_result_t ValidateExtractSubArrayQCOM(ValidationState_t& _, |
1103 | 0 | const Instruction* inst) { |
1104 | 0 | const auto result_type_inst = _.FindDef(inst->type_id()); |
1105 | 0 | const auto source = _.FindDef(inst->GetOperandAs<uint32_t>(2u)); |
1106 | 0 | const auto source_type_inst = _.FindDef(source->type_id()); |
1107 | | |
1108 | | // Are the input and the result arrays? |
1109 | 0 | if (result_type_inst->opcode() != spv::Op::OpTypeArray || |
1110 | 0 | source_type_inst->opcode() != spv::Op::OpTypeArray) { |
1111 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
1112 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
1113 | 0 | << " requires OpTypeArray operands for the input and the result."; |
1114 | 0 | } |
1115 | | |
1116 | 0 | const auto source_elt_type = _.GetComponentType(source_type_inst->id()); |
1117 | 0 | const auto result_elt_type = _.GetComponentType(result_type_inst->id()); |
1118 | | |
1119 | | // Do the input and result element types match? |
1120 | 0 | if (source_elt_type != result_elt_type) { |
1121 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
1122 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
1123 | 0 | << " requires the input and result element types match."; |
1124 | 0 | } |
1125 | | |
1126 | | // Elt type must be one of int32_t/uint32_t/float32/float16 |
1127 | 0 | if (!_.IsIntNOrFP32OrFP16<32>(source_elt_type)) { |
1128 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
1129 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
1130 | 0 | << " requires the element type be one of 32-bit OpTypeInt " |
1131 | 0 | "(signed/unsigned), 32-bit OpTypeFloat and 16-bit OpTypeFloat"; |
1132 | 0 | } |
1133 | | |
1134 | 0 | const auto start_index = _.FindDef(inst->GetOperandAs<uint32_t>(3u)); |
1135 | 0 | if (!start_index || !_.ContainsSizedIntOrFloatType(start_index->type_id(), |
1136 | 0 | spv::Op::OpTypeInt, 32)) { |
1137 | 0 | return _.diag(SPV_ERROR_INVALID_DATA, inst) |
1138 | 0 | << "Opcode " << spvOpcodeString(inst->opcode()) |
1139 | 0 | << " requires the type of the start index operand be 32-bit " |
1140 | 0 | "OpTypeInt"; |
1141 | 0 | } |
1142 | | |
1143 | 0 | return SPV_SUCCESS; |
1144 | 0 | } |
1145 | | |
1146 | | } // anonymous namespace |
1147 | | // Validates correctness of composite instructions. |
1148 | 14.3M | spv_result_t CompositesPass(ValidationState_t& _, const Instruction* inst) { |
1149 | 14.3M | switch (inst->opcode()) { |
1150 | 137 | case spv::Op::OpVectorExtractDynamic: |
1151 | 137 | return ValidateVectorExtractDynamic(_, inst); |
1152 | 45 | case spv::Op::OpVectorInsertDynamic: |
1153 | 45 | return ValidateVectorInsertDyanmic(_, inst); |
1154 | 14.8k | case spv::Op::OpVectorShuffle: |
1155 | 14.8k | return ValidateVectorShuffle(_, inst); |
1156 | 34.1k | case spv::Op::OpCompositeConstruct: |
1157 | 34.1k | return ValidateCompositeConstruct(_, inst); |
1158 | 0 | case spv::Op::OpCompositeConstructReplicateEXT: |
1159 | 0 | return ValidateCompositeConstructReplicate(_, inst); |
1160 | 50.9k | case spv::Op::OpCompositeExtract: |
1161 | 50.9k | return ValidateCompositeExtract(_, inst); |
1162 | 10.7k | case spv::Op::OpCompositeInsert: |
1163 | 10.7k | return ValidateCompositeInsert(_, inst); |
1164 | 1.17k | case spv::Op::OpCopyObject: |
1165 | 1.17k | return ValidateCopyObject(_, inst); |
1166 | 63 | case spv::Op::OpTranspose: |
1167 | 63 | return ValidateTranspose(_, inst); |
1168 | 0 | case spv::Op::OpCopyLogical: |
1169 | 0 | return ValidateCopyLogical(_, inst); |
1170 | 0 | case spv::Op::OpCompositeConstructCoopMatQCOM: |
1171 | 0 | return ValidateCompositeConstructCoopMatQCOM(_, inst); |
1172 | 0 | case spv::Op::OpCompositeExtractCoopMatQCOM: |
1173 | 0 | return ValidateCompositeExtractCoopMatQCOM(_, inst); |
1174 | 0 | case spv::Op::OpExtractSubArrayQCOM: |
1175 | 0 | return ValidateExtractSubArrayQCOM(_, inst); |
1176 | | |
1177 | 369 | case spv::Op::OpSpecConstantOp: { |
1178 | 369 | switch (inst->GetOperandAs<spv::Op>(2u)) { |
1179 | 7 | case spv::Op::OpVectorShuffle: |
1180 | 7 | return ValidateVectorShuffle(_, inst, 3); |
1181 | 3 | case spv::Op::OpCompositeExtract: |
1182 | 3 | return ValidateCompositeExtract(_, inst, 3); |
1183 | 2 | case spv::Op::OpCompositeInsert: |
1184 | 2 | return ValidateCompositeInsert(_, inst, 3); |
1185 | 357 | default: |
1186 | 357 | break; |
1187 | 369 | } |
1188 | 369 | } |
1189 | | |
1190 | 14.2M | default: |
1191 | 14.2M | break; |
1192 | 14.3M | } |
1193 | | |
1194 | 14.2M | return SPV_SUCCESS; |
1195 | 14.3M | } |
1196 | | |
1197 | | } // namespace val |
1198 | | } // namespace spvtools |