Coverage Report

Created: 2026-09-28 07:06

next uncovered line (L), next uncovered region (R), next uncovered branch (B)
/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