Coverage Report

Created: 2026-09-28 07:06

next uncovered line (L), next uncovered region (R), next uncovered branch (B)
/src/spirv-tools/source/opt/fold.cpp
Line
Count
Source
1
// Copyright (c) 2017 Google Inc.
2
//
3
// Licensed under the Apache License, Version 2.0 (the "License");
4
// you may not use this file except in compliance with the License.
5
// You may obtain a copy of the License at
6
//
7
//     http://www.apache.org/licenses/LICENSE-2.0
8
//
9
// Unless required by applicable law or agreed to in writing, software
10
// distributed under the License is distributed on an "AS IS" BASIS,
11
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
// See the License for the specific language governing permissions and
13
// limitations under the License.
14
15
#include "source/opt/fold.h"
16
17
#include <cassert>
18
#include <cstdint>
19
#include <vector>
20
21
#include "source/opt/const_folding_rules.h"
22
#include "source/opt/def_use_manager.h"
23
#include "source/opt/folding_rules.h"
24
#include "source/opt/ir_context.h"
25
26
namespace spvtools {
27
namespace opt {
28
namespace {
29
30
#ifndef INT32_MIN
31
#define INT32_MIN (-2147483648)
32
#endif
33
34
#ifndef INT32_MAX
35
#define INT32_MAX 2147483647
36
#endif
37
38
#ifndef UINT32_MAX
39
#define UINT32_MAX 0xffffffff /* 4294967295U */
40
#endif
41
42
}  // namespace
43
44
uint32_t InstructionFolder::UnaryOperate(spv::Op opcode,
45
2.35k
                                         uint32_t operand) const {
46
2.35k
  switch (opcode) {
47
    // Arthimetics
48
0
    case spv::Op::OpSNegate: {
49
0
      int32_t s_operand = static_cast<int32_t>(operand);
50
0
      if (s_operand == std::numeric_limits<int32_t>::min()) {
51
0
        return s_operand;
52
0
      }
53
0
      return static_cast<uint32_t>(-s_operand);
54
0
    }
55
1.51k
    case spv::Op::OpNot:
56
1.51k
      return ~operand;
57
837
    case spv::Op::OpLogicalNot:
58
837
      return !static_cast<bool>(operand);
59
0
    case spv::Op::OpUConvert:
60
0
      return operand;
61
0
    case spv::Op::OpSConvert:
62
0
      return operand;
63
0
    default:
64
0
      assert(false &&
65
0
             "Unsupported unary operation for OpSpecConstantOp instruction");
66
0
      return 0u;
67
2.35k
  }
68
2.35k
}
69
70
uint32_t InstructionFolder::BinaryOperate(spv::Op opcode, uint32_t a,
71
520k
                                          uint32_t b) const {
72
520k
  switch (opcode) {
73
    // Shifting
74
27.3k
    case spv::Op::OpShiftRightLogical:
75
27.3k
      if (b >= 32) {
76
        // This is undefined behaviour when |b| > 32.  Choose 0 for consistency.
77
        // When |b| == 32, doing the shift in C++ in undefined, but the result
78
        // will be 0, so just return that value.
79
23.9k
        return 0;
80
23.9k
      }
81
3.35k
      return a >> b;
82
11.8k
    case spv::Op::OpShiftRightArithmetic:
83
11.8k
      if (b > 32) {
84
        // This is undefined behaviour.  Choose 0 for consistency.
85
4.01k
        return 0;
86
4.01k
      }
87
7.82k
      if (b == 32) {
88
        // Doing the shift in C++ is undefined, but the result is defined in the
89
        // spir-v spec.  Find that value another way.
90
374
        if (static_cast<int32_t>(a) >= 0) {
91
327
          return 0;
92
327
        } else {
93
47
          return static_cast<uint32_t>(-1);
94
47
        }
95
374
      }
96
7.44k
      return (static_cast<int32_t>(a)) >> b;
97
8.84k
    case spv::Op::OpShiftLeftLogical:
98
8.84k
      if (b >= 32) {
99
        // This is undefined behaviour when |b| > 32.  Choose 0 for consistency.
100
        // When |b| == 32, doing the shift in C++ in undefined, but the result
101
        // will be 0, so just return that value.
102
3.53k
        return 0;
103
3.53k
      }
104
5.31k
      return a << b;
105
106
    // Bitwise operations
107
14.9k
    case spv::Op::OpBitwiseOr:
108
14.9k
      return a | b;
109
23.5k
    case spv::Op::OpBitwiseAnd:
110
23.5k
      return a & b;
111
27.3k
    case spv::Op::OpBitwiseXor:
112
27.3k
      return a ^ b;
113
114
    // Logical
115
402
    case spv::Op::OpLogicalEqual:
116
402
      return (static_cast<bool>(a)) == (static_cast<bool>(b));
117
1.09k
    case spv::Op::OpLogicalNotEqual:
118
1.09k
      return (static_cast<bool>(a)) != (static_cast<bool>(b));
119
559
    case spv::Op::OpLogicalOr:
120
559
      return (static_cast<bool>(a)) || (static_cast<bool>(b));
121
4.16k
    case spv::Op::OpLogicalAnd:
122
4.16k
      return (static_cast<bool>(a)) && (static_cast<bool>(b));
123
124
    // Comparison
125
26.9k
    case spv::Op::OpIEqual:
126
26.9k
      return a == b;
127
8.25k
    case spv::Op::OpINotEqual:
128
8.25k
      return a != b;
129
29.9k
    case spv::Op::OpULessThan:
130
29.9k
      return a < b;
131
145k
    case spv::Op::OpSLessThan:
132
145k
      return (static_cast<int32_t>(a)) < (static_cast<int32_t>(b));
133
15.4k
    case spv::Op::OpUGreaterThan:
134
15.4k
      return a > b;
135
32.5k
    case spv::Op::OpSGreaterThan:
136
32.5k
      return (static_cast<int32_t>(a)) > (static_cast<int32_t>(b));
137
7.54k
    case spv::Op::OpULessThanEqual:
138
7.54k
      return a <= b;
139
102k
    case spv::Op::OpSLessThanEqual:
140
102k
      return (static_cast<int32_t>(a)) <= (static_cast<int32_t>(b));
141
9.85k
    case spv::Op::OpUGreaterThanEqual:
142
9.85k
      return a >= b;
143
22.2k
    case spv::Op::OpSGreaterThanEqual:
144
22.2k
      return (static_cast<int32_t>(a)) >= (static_cast<int32_t>(b));
145
0
    default:
146
0
      assert(false &&
147
0
             "Unsupported binary operation for OpSpecConstantOp instruction");
148
0
      return 0u;
149
520k
  }
150
520k
}
151
152
uint32_t InstructionFolder::TernaryOperate(spv::Op opcode, uint32_t a,
153
2.95k
                                           uint32_t b, uint32_t c) const {
154
2.95k
  switch (opcode) {
155
2.95k
    case spv::Op::OpSelect:
156
2.95k
      return (static_cast<bool>(a)) ? b : c;
157
0
    default:
158
0
      assert(false &&
159
0
             "Unsupported ternary operation for OpSpecConstantOp instruction");
160
0
      return 0u;
161
2.95k
  }
162
2.95k
}
163
164
uint32_t InstructionFolder::OperateWords(
165
526k
    spv::Op opcode, const std::vector<uint32_t>& operand_words) const {
166
526k
  switch (operand_words.size()) {
167
2.35k
    case 1:
168
2.35k
      return UnaryOperate(opcode, operand_words.front());
169
520k
    case 2:
170
520k
      return BinaryOperate(opcode, operand_words.front(), operand_words.back());
171
2.95k
    case 3:
172
2.95k
      return TernaryOperate(opcode, operand_words[0], operand_words[1],
173
2.95k
                            operand_words[2]);
174
0
    default:
175
0
      assert(false && "Invalid number of operands");
176
0
      return 0;
177
526k
  }
178
526k
}
179
180
12.8M
bool InstructionFolder::FoldInstructionInternal(Instruction* inst) const {
181
12.8M
  auto identity_map = [](uint32_t id) { return id; };
182
12.8M
  Instruction* folded_inst = FoldInstructionToConstant(inst, identity_map);
183
12.8M
  if (folded_inst != nullptr) {
184
1.95M
    inst->SetOpcode(spv::Op::OpCopyObject);
185
1.95M
    inst->SetInOperands({{SPV_OPERAND_TYPE_ID, {folded_inst->result_id()}}});
186
1.95M
    return true;
187
1.95M
  }
188
189
10.9M
  analysis::ConstantManager* const_manager = context_->get_constant_mgr();
190
10.9M
  std::vector<const analysis::Constant*> constants =
191
10.9M
      const_manager->GetOperandConstants(inst);
192
193
10.9M
  for (const FoldingRule& rule :
194
12.3M
       GetFoldingRules().GetRulesForInstruction(inst)) {
195
12.3M
    if (rule(context_, inst, constants)) {
196
370k
      return true;
197
370k
    }
198
12.3M
  }
199
10.5M
  return false;
200
10.9M
}
201
202
// Returns the result of performing an operation on scalar constant operands.
203
// This function extracts the operand values as 32 bit words and returns the
204
// result in 32 bit word. Scalar constants with longer than 32-bit width are
205
// not accepted in this function.
206
uint32_t InstructionFolder::FoldScalars(
207
    spv::Op opcode,
208
524k
    const std::vector<const analysis::Constant*>& operands) const {
209
524k
  assert(IsFoldableOpcode(opcode) &&
210
524k
         "Unhandled instruction opcode in FoldScalars");
211
524k
  std::vector<uint32_t> operand_values_in_raw_words;
212
1.04M
  for (const auto& operand : operands) {
213
1.04M
    if (const analysis::ScalarConstant* scalar = operand->AsScalarConstant()) {
214
1.04M
      const auto& scalar_words = scalar->words();
215
1.04M
      assert(scalar_words.size() == 1 &&
216
1.04M
             "Scalar constants with longer than 32-bit width are not allowed "
217
1.04M
             "in FoldScalars()");
218
1.04M
      operand_values_in_raw_words.push_back(scalar_words.front());
219
1.04M
    } else if (operand->AsNullConstant()) {
220
622
      operand_values_in_raw_words.push_back(0u);
221
622
    } else {
222
0
      assert(false &&
223
0
             "FoldScalars() only accepts ScalarConst or NullConst type of "
224
0
             "constant");
225
0
    }
226
1.04M
  }
227
524k
  return OperateWords(opcode, operand_values_in_raw_words);
228
524k
}
229
230
bool InstructionFolder::FoldBinaryIntegerOpToConstant(
231
    Instruction* inst, const std::function<uint32_t(uint32_t)>& id_map,
232
632k
    uint32_t* result) const {
233
632k
  spv::Op opcode = inst->opcode();
234
632k
  analysis::ConstantManager* const_manger = context_->get_constant_mgr();
235
236
632k
  uint32_t ids[2];
237
632k
  const analysis::IntConstant* constants[2];
238
1.89M
  for (uint32_t i = 0; i < 2; i++) {
239
1.26M
    const Operand* operand = &inst->GetInOperand(i);
240
1.26M
    if (operand->type != SPV_OPERAND_TYPE_ID) {
241
0
      return false;
242
0
    }
243
1.26M
    ids[i] = id_map(operand->words[0]);
244
1.26M
    const analysis::Constant* constant =
245
1.26M
        const_manger->FindDeclaredConstant(ids[i]);
246
1.26M
    constants[i] = (constant != nullptr ? constant->AsIntConstant() : nullptr);
247
1.26M
  }
248
249
632k
  switch (opcode) {
250
    // Arthimetics
251
21.9k
    case spv::Op::OpIMul:
252
64.6k
      for (uint32_t i = 0; i < 2; i++) {
253
43.6k
        if (constants[i] != nullptr && constants[i]->IsZero()) {
254
917
          *result = 0;
255
917
          return true;
256
917
        }
257
43.6k
      }
258
21.0k
      break;
259
21.0k
    case spv::Op::OpUDiv:
260
21.2k
    case spv::Op::OpSDiv:
261
42.3k
    case spv::Op::OpSRem:
262
52.1k
    case spv::Op::OpSMod:
263
52.4k
    case spv::Op::OpUMod:
264
      // This changes undefined behaviour (ie divide by 0) into a 0.
265
154k
      for (uint32_t i = 0; i < 2; i++) {
266
104k
        if (constants[i] != nullptr && constants[i]->IsZero()) {
267
2.38k
          *result = 0;
268
2.38k
          return true;
269
2.38k
        }
270
104k
      }
271
50.0k
      break;
272
273
    // Shifting
274
50.0k
    case spv::Op::OpShiftRightLogical:
275
11.8k
    case spv::Op::OpShiftLeftLogical:
276
11.8k
      if (constants[1] != nullptr) {
277
        // When shifting by a value larger than the size of the result, the
278
        // result is undefined.  We are setting the undefined behaviour to a
279
        // result of 0.  If the shift amount is the same as the size of the
280
        // result, then the result is defined, and it 0.
281
8.56k
        uint32_t shift_amount = constants[1]->GetU32BitValue();
282
8.56k
        if (shift_amount >= 32) {
283
4.12k
          *result = 0;
284
4.12k
          return true;
285
4.12k
        }
286
8.56k
      }
287
7.72k
      break;
288
289
    // Bitwise operations
290
75.4k
    case spv::Op::OpBitwiseOr:
291
197k
      for (uint32_t i = 0; i < 2; i++) {
292
138k
        if (constants[i] != nullptr) {
293
          // TODO: Change the mask against a value based on the bit width of the
294
          // instruction result type.  This way we can handle say 16-bit values
295
          // as well.
296
32.9k
          uint32_t mask = constants[i]->GetU32BitValue();
297
32.9k
          if (mask == 0xFFFFFFFF) {
298
15.8k
            *result = 0xFFFFFFFF;
299
15.8k
            return true;
300
15.8k
          }
301
32.9k
        }
302
138k
      }
303
59.5k
      break;
304
59.5k
    case spv::Op::OpBitwiseAnd:
305
70.2k
      for (uint32_t i = 0; i < 2; i++) {
306
47.4k
        if (constants[i] != nullptr) {
307
17.5k
          if (constants[i]->IsZero()) {
308
1.20k
            *result = 0;
309
1.20k
            return true;
310
1.20k
          }
311
17.5k
        }
312
47.4k
      }
313
22.7k
      break;
314
315
    // Comparison
316
22.7k
    case spv::Op::OpULessThan:
317
5.87k
      if (constants[0] != nullptr &&
318
1.37k
          constants[0]->GetU32BitValue() == UINT32_MAX) {
319
34
        *result = false;
320
34
        return true;
321
34
      }
322
5.84k
      if (constants[1] != nullptr && constants[1]->GetU32BitValue() == 0) {
323
490
        *result = false;
324
490
        return true;
325
490
      }
326
5.35k
      break;
327
45.1k
    case spv::Op::OpSLessThan:
328
45.1k
      if (constants[0] != nullptr &&
329
1.07k
          constants[0]->GetS32BitValue() == INT32_MAX) {
330
4
        *result = false;
331
4
        return true;
332
4
      }
333
45.1k
      if (constants[1] != nullptr &&
334
34.2k
          constants[1]->GetS32BitValue() == INT32_MIN) {
335
20
        *result = false;
336
20
        return true;
337
20
      }
338
45.1k
      break;
339
45.1k
    case spv::Op::OpUGreaterThan:
340
9.55k
      if (constants[0] != nullptr && constants[0]->IsZero()) {
341
403
        *result = false;
342
403
        return true;
343
403
      }
344
9.15k
      if (constants[1] != nullptr &&
345
7.46k
          constants[1]->GetU32BitValue() == UINT32_MAX) {
346
74
        *result = false;
347
74
        return true;
348
74
      }
349
9.07k
      break;
350
31.0k
    case spv::Op::OpSGreaterThan:
351
31.0k
      if (constants[0] != nullptr &&
352
7.10k
          constants[0]->GetS32BitValue() == INT32_MIN) {
353
17
        *result = false;
354
17
        return true;
355
17
      }
356
31.0k
      if (constants[1] != nullptr &&
357
14.9k
          constants[1]->GetS32BitValue() == INT32_MAX) {
358
5
        *result = false;
359
5
        return true;
360
5
      }
361
31.0k
      break;
362
31.0k
    case spv::Op::OpULessThanEqual:
363
12.2k
      if (constants[0] != nullptr && constants[0]->IsZero()) {
364
424
        *result = true;
365
424
        return true;
366
424
      }
367
11.8k
      if (constants[1] != nullptr &&
368
6.27k
          constants[1]->GetU32BitValue() == UINT32_MAX) {
369
122
        *result = true;
370
122
        return true;
371
122
      }
372
11.7k
      break;
373
44.7k
    case spv::Op::OpSLessThanEqual:
374
44.7k
      if (constants[0] != nullptr &&
375
12.0k
          constants[0]->GetS32BitValue() == INT32_MIN) {
376
45
        *result = true;
377
45
        return true;
378
45
      }
379
44.6k
      if (constants[1] != nullptr &&
380
14.3k
          constants[1]->GetS32BitValue() == INT32_MAX) {
381
33
        *result = true;
382
33
        return true;
383
33
      }
384
44.6k
      break;
385
44.6k
    case spv::Op::OpUGreaterThanEqual:
386
7.22k
      if (constants[0] != nullptr &&
387
1.74k
          constants[0]->GetU32BitValue() == UINT32_MAX) {
388
35
        *result = true;
389
35
        return true;
390
35
      }
391
7.19k
      if (constants[1] != nullptr && constants[1]->GetU32BitValue() == 0) {
392
1.00k
        *result = true;
393
1.00k
        return true;
394
1.00k
      }
395
6.18k
      break;
396
11.6k
    case spv::Op::OpSGreaterThanEqual:
397
11.6k
      if (constants[0] != nullptr &&
398
1.37k
          constants[0]->GetS32BitValue() == INT32_MAX) {
399
4
        *result = true;
400
4
        return true;
401
4
      }
402
11.6k
      if (constants[1] != nullptr &&
403
7.05k
          constants[1]->GetS32BitValue() == INT32_MIN) {
404
18
        *result = true;
405
18
        return true;
406
18
      }
407
11.6k
      break;
408
279k
    default:
409
279k
      break;
410
632k
  }
411
605k
  return false;
412
632k
}
413
414
bool InstructionFolder::FoldBinaryBooleanOpToConstant(
415
    Instruction* inst, const std::function<uint32_t(uint32_t)>& id_map,
416
605k
    uint32_t* result) const {
417
605k
  spv::Op opcode = inst->opcode();
418
605k
  analysis::ConstantManager* const_manger = context_->get_constant_mgr();
419
420
605k
  uint32_t ids[2];
421
605k
  const analysis::BoolConstant* constants[2];
422
1.81M
  for (uint32_t i = 0; i < 2; i++) {
423
1.21M
    const Operand* operand = &inst->GetInOperand(i);
424
1.21M
    if (operand->type != SPV_OPERAND_TYPE_ID) {
425
0
      return false;
426
0
    }
427
1.21M
    ids[i] = id_map(operand->words[0]);
428
1.21M
    const analysis::Constant* constant =
429
1.21M
        const_manger->FindDeclaredConstant(ids[i]);
430
1.21M
    constants[i] = (constant != nullptr ? constant->AsBoolConstant() : nullptr);
431
1.21M
  }
432
433
605k
  switch (opcode) {
434
    // Logical
435
1.64k
    case spv::Op::OpLogicalOr:
436
4.11k
      for (uint32_t i = 0; i < 2; i++) {
437
3.20k
        if (constants[i] != nullptr) {
438
997
          if (constants[i]->value()) {
439
745
            *result = true;
440
745
            return true;
441
745
          }
442
997
        }
443
3.20k
      }
444
903
      break;
445
11.0k
    case spv::Op::OpLogicalAnd:
446
32.5k
      for (uint32_t i = 0; i < 2; i++) {
447
21.8k
        if (constants[i] != nullptr) {
448
5.80k
          if (!constants[i]->value()) {
449
376
            *result = false;
450
376
            return true;
451
376
          }
452
5.80k
        }
453
21.8k
      }
454
10.6k
      break;
455
456
592k
    default:
457
592k
      break;
458
605k
  }
459
604k
  return false;
460
605k
}
461
462
bool InstructionFolder::FoldIntegerOpToConstant(
463
    Instruction* inst, const std::function<uint32_t(uint32_t)>& id_map,
464
684k
    uint32_t* result) const {
465
684k
  assert(IsFoldableOpcode(inst->opcode()) &&
466
684k
         "Unhandled instruction opcode in FoldScalars");
467
684k
  switch (inst->NumInOperands()) {
468
632k
    case 2:
469
632k
      return FoldBinaryIntegerOpToConstant(inst, id_map, result) ||
470
605k
             FoldBinaryBooleanOpToConstant(inst, id_map, result);
471
52.3k
    default:
472
52.3k
      return false;
473
684k
  }
474
684k
}
475
476
std::vector<uint32_t> InstructionFolder::FoldVectors(
477
    spv::Op opcode, uint32_t num_dims,
478
523
    const std::vector<const analysis::Constant*>& operands) const {
479
523
  assert(IsFoldableOpcode(opcode) &&
480
523
         "Unhandled instruction opcode in FoldVectors");
481
523
  std::vector<uint32_t> result;
482
1.76k
  for (uint32_t d = 0; d < num_dims; d++) {
483
1.24k
    std::vector<uint32_t> operand_values_for_one_dimension;
484
3.28k
    for (const auto& operand : operands) {
485
3.28k
      if (const analysis::VectorConstant* vector_operand =
486
3.28k
              operand->AsVectorConstant()) {
487
        // Extract the raw value of the scalar component constants
488
        // in 32-bit words here. The reason of not using FoldScalars() here
489
        // is that we do not create temporary null constants as components
490
        // when the vector operand is a NullConstant because Constant creation
491
        // may need extra checks for the validity and that is not managed in
492
        // here.
493
3.28k
        if (const analysis::ScalarConstant* scalar_component =
494
3.28k
                vector_operand->GetComponents().at(d)->AsScalarConstant()) {
495
3.28k
          const auto& scalar_words = scalar_component->words();
496
3.28k
          assert(
497
3.28k
              scalar_words.size() == 1 &&
498
3.28k
              "Vector components with longer than 32-bit width are not allowed "
499
3.28k
              "in FoldVectors()");
500
3.28k
          operand_values_for_one_dimension.push_back(scalar_words.front());
501
3.28k
        } else if (operand->AsNullConstant()) {
502
0
          operand_values_for_one_dimension.push_back(0u);
503
0
        } else {
504
0
          assert(false &&
505
0
                 "VectorConst should only has ScalarConst or NullConst as "
506
0
                 "components");
507
0
        }
508
3.28k
      } else if (operand->AsNullConstant()) {
509
0
        operand_values_for_one_dimension.push_back(0u);
510
0
      } else {
511
0
        assert(false &&
512
0
               "FoldVectors() only accepts VectorConst or NullConst type of "
513
0
               "constant");
514
0
      }
515
3.28k
    }
516
1.24k
    result.push_back(OperateWords(opcode, operand_values_for_one_dimension));
517
1.24k
  }
518
523
  return result;
519
523
}
520
521
34.7M
bool InstructionFolder::IsFoldableOpcode(spv::Op opcode) const {
522
  // NOTE: Extend to more opcodes as new cases are handled in the folder
523
  // functions.
524
34.7M
  switch (opcode) {
525
152k
    case spv::Op::OpBitwiseAnd:
526
458k
    case spv::Op::OpBitwiseOr:
527
550k
    case spv::Op::OpBitwiseXor:
528
1.32M
    case spv::Op::OpIAdd:
529
1.59M
    case spv::Op::OpIEqual:
530
1.72M
    case spv::Op::OpIMul:
531
1.80M
    case spv::Op::OpINotEqual:
532
2.23M
    case spv::Op::OpISub:
533
2.28M
    case spv::Op::OpLogicalAnd:
534
2.28M
    case spv::Op::OpLogicalEqual:
535
2.37M
    case spv::Op::OpLogicalNot:
536
2.38M
    case spv::Op::OpLogicalNotEqual:
537
2.39M
    case spv::Op::OpLogicalOr:
538
2.42M
    case spv::Op::OpNot:
539
2.53M
    case spv::Op::OpSDiv:
540
2.61M
    case spv::Op::OpSelect:
541
2.81M
    case spv::Op::OpSGreaterThan:
542
2.91M
    case spv::Op::OpSGreaterThanEqual:
543
2.96M
    case spv::Op::OpShiftLeftLogical:
544
3.00M
    case spv::Op::OpShiftRightArithmetic:
545
3.10M
    case spv::Op::OpShiftRightLogical:
546
3.69M
    case spv::Op::OpSLessThan:
547
4.16M
    case spv::Op::OpSLessThanEqual:
548
4.19M
    case spv::Op::OpSMod:
549
4.21M
    case spv::Op::OpSNegate:
550
4.34M
    case spv::Op::OpSRem:
551
4.34M
    case spv::Op::OpSConvert:
552
4.34M
    case spv::Op::OpUConvert:
553
4.34M
    case spv::Op::OpUDiv:
554
4.42M
    case spv::Op::OpUGreaterThan:
555
4.48M
    case spv::Op::OpUGreaterThanEqual:
556
4.58M
    case spv::Op::OpULessThan:
557
4.65M
    case spv::Op::OpULessThanEqual:
558
4.65M
    case spv::Op::OpUMod:
559
4.65M
      return true;
560
30.1M
    default:
561
30.1M
      return false;
562
34.7M
  }
563
34.7M
}
564
565
bool InstructionFolder::IsFoldableConstant(
566
0
    const analysis::Constant* cst) const {
567
  // Currently supported constants are 32-bit values or null constants.
568
0
  if (const analysis::ScalarConstant* scalar = cst->AsScalarConstant())
569
0
    return scalar->words().size() == 1;
570
0
  else
571
0
    return cst->AsNullConstant() != nullptr;
572
0
}
573
574
Instruction* InstructionFolder::FoldInstructionToConstant(
575
13.5M
    Instruction* inst, std::function<uint32_t(uint32_t)> id_map) const {
576
13.5M
  analysis::ConstantManager* const_mgr = context_->get_constant_mgr();
577
578
13.5M
  if (!inst->IsFoldableByFoldScalar() && !inst->IsFoldableByFoldVector() &&
579
11.5M
      !GetConstantFoldingRules().HasFoldingRule(inst)) {
580
8.41M
    return nullptr;
581
8.41M
  }
582
  // Collect the values of the constant parameters.
583
5.08M
  std::vector<const analysis::Constant*> constants;
584
5.08M
  bool missing_constants = false;
585
5.08M
  inst->ForEachInId([&constants, &missing_constants, const_mgr,
586
9.59M
                     &id_map](uint32_t* op_id) {
587
9.59M
    uint32_t id = id_map(*op_id);
588
9.59M
    const analysis::Constant* const_op = const_mgr->FindDeclaredConstant(id);
589
9.59M
    if (!const_op) {
590
4.51M
      constants.push_back(nullptr);
591
4.51M
      missing_constants = true;
592
5.07M
    } else {
593
5.07M
      constants.push_back(const_op);
594
5.07M
    }
595
9.59M
  });
596
597
5.08M
  const analysis::Constant* folded_const = nullptr;
598
5.08M
  for (auto rule : GetConstantFoldingRules().GetRulesForInstruction(inst)) {
599
4.63M
    folded_const = rule(context_, inst, constants);
600
4.63M
    if (folded_const == nullptr && inst->context()->id_overflow()) {
601
0
      return nullptr;
602
0
    }
603
4.63M
    if (folded_const != nullptr) {
604
1.55M
      Instruction* const_inst =
605
1.55M
          const_mgr->GetDefiningInstruction(folded_const, inst->type_id());
606
1.55M
      if (const_inst == nullptr) {
607
2
        return nullptr;
608
2
      }
609
1.55M
      assert(const_inst->type_id() == inst->type_id());
610
      // May be a new instruction that needs to be analysed.
611
1.55M
      context_->UpdateDefUse(const_inst);
612
1.55M
      return const_inst;
613
1.55M
    }
614
4.63M
  }
615
616
3.53M
  bool successful = false;
617
618
  // If all parameters are constant, fold the instruction to a constant.
619
3.53M
  if (inst->IsFoldableByFoldScalar()) {
620
1.20M
    uint32_t result_val = 0;
621
622
1.20M
    if (!missing_constants) {
623
524k
      result_val = FoldScalars(inst->opcode(), constants);
624
524k
      successful = true;
625
524k
    }
626
627
1.20M
    if (!successful) {
628
684k
      successful = FoldIntegerOpToConstant(inst, id_map, &result_val);
629
684k
    }
630
631
1.20M
    if (successful) {
632
553k
      const analysis::Constant* result_const =
633
553k
          const_mgr->GetConstant(const_mgr->GetType(inst), {result_val});
634
553k
      if (!result_const) {
635
0
        return nullptr;
636
0
      }
637
553k
      Instruction* folded_inst =
638
553k
          const_mgr->GetDefiningInstruction(result_const, inst->type_id());
639
553k
      return folded_inst;
640
553k
    }
641
2.32M
  } else if (inst->IsFoldableByFoldVector()) {
642
2.16k
    std::vector<uint32_t> result_val;
643
644
2.16k
    if (!missing_constants) {
645
523
      if (Instruction* inst_type =
646
523
              context_->get_def_use_mgr()->GetDef(inst->type_id())) {
647
523
        result_val = FoldVectors(
648
523
            inst->opcode(), inst_type->GetSingleWordInOperand(1), constants);
649
523
        successful = true;
650
523
      }
651
523
    }
652
653
2.16k
    if (successful) {
654
523
      const analysis::Constant* result_const =
655
523
          const_mgr->GetNumericVectorConstantWithWords(
656
523
              const_mgr->GetType(inst)->AsVector(), result_val);
657
523
      if (!result_const) {
658
0
        return nullptr;
659
0
      }
660
523
      Instruction* folded_inst =
661
523
          const_mgr->GetDefiningInstruction(result_const, inst->type_id());
662
523
      return folded_inst;
663
523
    }
664
2.16k
  }
665
666
2.98M
  return nullptr;
667
3.53M
}
668
669
0
bool InstructionFolder::IsFoldableType(Instruction* type_inst) const {
670
0
  return IsFoldableScalarType(type_inst) || IsFoldableVectorType(type_inst);
671
0
}
672
673
10.2M
bool InstructionFolder::IsFoldableScalarType(Instruction* type_inst) const {
674
  // Support 32-bit integers.
675
10.2M
  if (type_inst->opcode() == spv::Op::OpTypeInt) {
676
8.47M
    return type_inst->GetSingleWordInOperand(0) == 32;
677
8.47M
  }
678
  // Support booleans.
679
1.73M
  if (type_inst->opcode() == spv::Op::OpTypeBool) {
680
1.71M
    return true;
681
1.71M
  }
682
  // Nothing else yet.
683
15.9k
  return false;
684
1.73M
}
685
686
30.5k
bool InstructionFolder::IsFoldableVectorType(Instruction* type_inst) const {
687
  // Support vectors with foldable components
688
30.5k
  if (type_inst->opcode() == spv::Op::OpTypeVector) {
689
20.8k
    uint32_t component_type_id = type_inst->GetSingleWordInOperand(0);
690
20.8k
    Instruction* def_component_type =
691
20.8k
        context_->get_def_use_mgr()->GetDef(component_type_id);
692
20.8k
    return def_component_type != nullptr &&
693
20.8k
           IsFoldableScalarType(def_component_type);
694
20.8k
  }
695
  // Nothing else yet.
696
9.74k
  return false;
697
30.5k
}
698
699
12.7M
bool InstructionFolder::FoldInstruction(Instruction* inst) const {
700
12.7M
  bool modified = false;
701
12.7M
  Instruction* folded_inst(inst);
702
15.0M
  while (folded_inst->opcode() != spv::Op::OpCopyObject &&
703
12.8M
         FoldInstructionInternal(&*folded_inst)) {
704
2.32M
    modified = true;
705
2.32M
  }
706
12.7M
  return modified;
707
12.7M
}
708
709
}  // namespace opt
710
}  // namespace spvtools