MLIR 24.0.0git
SPIRVCanonicalization.cpp
Go to the documentation of this file.
1//===- SPIRVCanonicalization.cpp - MLIR SPIR-V canonicalization patterns --===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file defines the folders and canonicalization patterns for SPIR-V ops.
10//
11//===----------------------------------------------------------------------===//
12
13#include <optional>
14#include <utility>
15
17
21#include "mlir/IR/Matchers.h"
23#include "llvm/ADT/STLExtras.h"
24#include "llvm/ADT/SmallVectorExtras.h"
25
26using namespace mlir;
27
28//===----------------------------------------------------------------------===//
29// Common utility functions
30//===----------------------------------------------------------------------===//
31
32/// Returns the boolean value under the hood if the given `boolAttr` is a scalar
33/// or splat vector bool constant.
34static std::optional<bool> getScalarOrSplatBoolAttr(Attribute attr) {
35 if (!attr)
36 return std::nullopt;
37
38 if (auto boolAttr = dyn_cast<BoolAttr>(attr))
39 return boolAttr.getValue();
40 if (auto splatAttr = dyn_cast<SplatElementsAttr>(attr))
41 if (splatAttr.getElementType().isInteger(1))
42 return splatAttr.getSplatValue<bool>();
43 return std::nullopt;
44}
45
46// Extracts an element from the given `composite` by following the given
47// `indices`. Returns a null Attribute if error happens.
50 // Check that given composite is a constant.
51 if (!composite)
52 return {};
53 // Return composite itself if we reach the end of the index chain.
54 if (indices.empty())
55 return composite;
56
57 if (auto vector = dyn_cast<ElementsAttr>(composite)) {
58 assert(indices.size() == 1 && "must have exactly one index for a vector");
59 return vector.getValues<Attribute>()[indices[0]];
60 }
61
62 if (auto array = dyn_cast<ArrayAttr>(composite)) {
63 assert(!indices.empty() && "must have at least one index for an array");
64 return extractCompositeElement(array.getValue()[indices[0]],
65 indices.drop_front());
66 }
67
68 return {};
69}
70
71static bool isDivZeroOrOverflow(const APInt &a, const APInt &b) {
72 bool div0 = b.isZero();
73 bool overflow = a.isMinSignedValue() && b.isAllOnes();
74
75 return div0 || overflow;
76}
77
78//===----------------------------------------------------------------------===//
79// TableGen'erated canonicalizers
80//===----------------------------------------------------------------------===//
81
82namespace {
83#include "SPIRVCanonicalization.inc"
84} // namespace
85
86//===----------------------------------------------------------------------===//
87// spirv.AccessChainOp / spirv.InBoundsAccessChainOp
88//===----------------------------------------------------------------------===//
89
90namespace {
91
92/// Combines chained SPIR-V access chain operations of the same kind into one.
93template <typename AccessChainOp>
94struct CombineChainedAccessChain final : OpRewritePattern<AccessChainOp> {
95 using OpRewritePattern<AccessChainOp>::OpRewritePattern;
96
97 LogicalResult matchAndRewrite(AccessChainOp accessChainOp,
98 PatternRewriter &rewriter) const override {
99 auto parentAccessChainOp =
100 accessChainOp.getBasePtr().template getDefiningOp<AccessChainOp>();
101
102 if (!parentAccessChainOp) {
103 return failure();
104 }
105
106 // Combine indices.
107 SmallVector<Value, 4> indices(parentAccessChainOp.getIndices());
108 llvm::append_range(indices, accessChainOp.getIndices());
109
110 rewriter.replaceOpWithNewOp<AccessChainOp>(
111 accessChainOp, parentAccessChainOp.getBasePtr(), indices);
112
113 return success();
114 }
115};
116} // namespace
117
118void spirv::AccessChainOp::getCanonicalizationPatterns(
119 RewritePatternSet &results, MLIRContext *context) {
120 results.add<CombineChainedAccessChain<spirv::AccessChainOp>>(context);
121}
122
123void spirv::InBoundsAccessChainOp::getCanonicalizationPatterns(
124 RewritePatternSet &results, MLIRContext *context) {
125 results.add<CombineChainedAccessChain<spirv::InBoundsAccessChainOp>>(context);
126}
127
128//===----------------------------------------------------------------------===//
129// spirv.IAddCarry / spirv.ISubBorrow
130//===----------------------------------------------------------------------===//
131
132template <typename Op>
135
136 static constexpr bool IsSub = std::is_same_v<Op, spirv::ISubBorrowOp>;
137
138 LogicalResult matchAndRewrite(Op op,
139 PatternRewriter &rewriter) const override {
140 Value lhs = op.getOperand1();
141 Value rhs = op.getOperand2();
142
143 // iaddcarry (x, 0) = <0, x>
144 // isubborrow (x, 0) = <x, 0>
145 if (matchPattern(rhs, m_Zero())) {
146 std::array<Value, 2> constituents =
147 IsSub ? std::array{lhs, rhs} : std::array{rhs, lhs};
148 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(op, op.getType(),
149 constituents);
150 return success();
151 }
152
153 Attribute lhsAttr;
154 Attribute rhsAttr;
155 if (!matchPattern(lhs, m_Constant(&lhsAttr)) ||
156 !matchPattern(rhs, m_Constant(&rhsAttr)))
157 return failure();
158
159 auto lowBits = constFoldBinaryOp<IntegerAttr>(
160 {lhsAttr, rhsAttr},
161 [](const APInt &a, const APInt &b) { return IsSub ? a - b : a + b; });
162 if (!lowBits)
163 return failure();
164
165 auto wrapBit = constFoldBinaryOp<IntegerAttr>(
166 {lhsAttr, rhsAttr}, [](const APInt &a, const APInt &b) {
167 bool wrapped = IsSub ? a.ult(b) : (a + b).ult(a);
168 return APInt(a.getBitWidth(), wrapped ? 1 : 0);
169 });
170 if (!wrapBit)
171 return failure();
172
173 rewriter.replaceOpWithNewOp<spirv::ConstantOp>(
174 op, op.getType(), rewriter.getArrayAttr({lowBits, wrapBit}));
175 return success();
176 }
177};
178
180void spirv::IAddCarryOp::getCanonicalizationPatterns(
181 RewritePatternSet &patterns, MLIRContext *context) {
182 patterns.add<IAddCarryFold>(context);
183}
184
186void spirv::ISubBorrowOp::getCanonicalizationPatterns(
187 RewritePatternSet &patterns, MLIRContext *context) {
188 patterns.add<ISubBorrowFold>(context);
189}
190
191//===----------------------------------------------------------------------===//
192// spirv.[S|U]MulExtended
193//===----------------------------------------------------------------------===//
194
195template <typename MulOp, bool IsSigned>
196struct MulExtendedFold final : OpRewritePattern<MulOp> {
198
199 LogicalResult matchAndRewrite(MulOp op,
200 PatternRewriter &rewriter) const override {
201 Location loc = op.getLoc();
202 Value lhs = op.getOperand1();
203 Value rhs = op.getOperand2();
204 Type constituentType = lhs.getType();
205
206 // [su]mulextended (x, 0) = <0, 0>
207 if (matchPattern(rhs, m_Zero())) {
208 Value zero = spirv::ConstantOp::getZero(constituentType, loc, rewriter);
209 Value constituents[2] = {zero, zero};
210 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(op, op.getType(),
211 constituents);
212 return success();
213 }
214
215 // According to the SPIR-V spec:
216 //
217 // Result Type must be from OpTypeStruct. The struct must have two
218 // members...
219 //
220 // Member 0 of the result gets the low-order bits of the multiplication.
221 //
222 // Member 1 of the result gets the high-order bits of the multiplication.
223 Attribute lhsAttr;
224 Attribute rhsAttr;
225 if (!matchPattern(lhs, m_Constant(&lhsAttr)) ||
226 !matchPattern(rhs, m_Constant(&rhsAttr)))
227 return failure();
228
229 auto lowBits = constFoldBinaryOp<IntegerAttr>(
230 {lhsAttr, rhsAttr},
231 [](const APInt &a, const APInt &b) { return a * b; });
232
233 if (!lowBits)
234 return failure();
235
236 auto highBits = constFoldBinaryOp<IntegerAttr>(
237 {lhsAttr, rhsAttr}, [](const APInt &a, const APInt &b) {
238 if (IsSigned) {
239 return llvm::APIntOps::mulhs(a, b);
240 }
241 return llvm::APIntOps::mulhu(a, b);
242 });
243
244 if (!highBits)
245 return failure();
246
247 rewriter.replaceOpWithNewOp<spirv::ConstantOp>(
248 op, op.getType(), rewriter.getArrayAttr({lowBits, highBits}));
249 return success();
250 }
251};
252
253template <typename MulOp>
254struct MulExtendedOpXOne final : OpRewritePattern<MulOp> {
256
257 LogicalResult matchAndRewrite(MulOp op,
258 PatternRewriter &rewriter) const override {
259 Location loc = op.getLoc();
260 Value lhs = op.getOperand1();
261 Value rhs = op.getOperand2();
262 Type constituentType = lhs.getType();
263
264 // [su]mulextended (x, 1) = <x, 0>
265 if (matchPattern(rhs, m_One())) {
266 Value zero = spirv::ConstantOp::getZero(constituentType, loc, rewriter);
267 Value constituents[2] = {lhs, zero};
268 rewriter.replaceOpWithNewOp<spirv::CompositeConstructOp>(op, op.getType(),
269 constituents);
270 return success();
271 }
272
273 return failure();
274 }
275};
276
279void spirv::SMulExtendedOp::getCanonicalizationPatterns(
280 RewritePatternSet &patterns, MLIRContext *context) {
281 patterns.add<SMulExtendedOpFold, SMulExtendedOpXOne>(context);
282}
283
286void spirv::UMulExtendedOp::getCanonicalizationPatterns(
287 RewritePatternSet &patterns, MLIRContext *context) {
288 patterns.add<UMulExtendedOpFold, UMulExtendedOpXOne>(context);
289}
290
291//===----------------------------------------------------------------------===//
292// spirv.UMod
293//===----------------------------------------------------------------------===//
294
295// Input:
296// %0 = spirv.UMod %arg0, %const32 : i32
297// %1 = spirv.UMod %0, %const4 : i32
298// Output:
299// %0 = spirv.UMod %arg0, %const32 : i32
300// %1 = spirv.UMod %arg0, %const4 : i32
301
302// The transformation is only applied if one divisor is a multiple of the other.
303
304struct UModSimplification final : OpRewritePattern<spirv::UModOp> {
305 using Base::Base;
306
307 LogicalResult matchAndRewrite(spirv::UModOp umodOp,
308 PatternRewriter &rewriter) const override {
309 auto prevUMod = umodOp.getOperand(0).getDefiningOp<spirv::UModOp>();
310 if (!prevUMod)
311 return failure();
312
313 TypedAttr prevValue;
314 TypedAttr currValue;
315 if (!matchPattern(prevUMod.getOperand(1), m_Constant(&prevValue)) ||
316 !matchPattern(umodOp.getOperand(1), m_Constant(&currValue)))
317 return failure();
318
319 // Ensure that previous divisor is a multiple of the current divisor. If
320 // not, fail the transformation.
321 bool isApplicable = false;
322 if (auto prevInt = dyn_cast<IntegerAttr>(prevValue)) {
323 auto currInt = cast<IntegerAttr>(currValue);
324 if (currInt.getValue().isZero())
325 return failure();
326 isApplicable = prevInt.getValue().urem(currInt.getValue()) == 0;
327 } else if (auto prevVec = dyn_cast<DenseElementsAttr>(prevValue)) {
328 auto currVec = cast<DenseElementsAttr>(currValue);
329 if (llvm::any_of(currVec.getValues<APInt>(),
330 [](const APInt &curr) { return curr.isZero(); }))
331 return failure();
332 isApplicable = llvm::all_of(llvm::zip_equal(prevVec.getValues<APInt>(),
333 currVec.getValues<APInt>()),
334 [](const auto &pair) {
335 auto &[prev, curr] = pair;
336 return prev.urem(curr) == 0;
337 });
338 }
339
340 if (!isApplicable)
341 return failure();
342
343 // The transformation is safe. Replace the existing UMod operation with a
344 // new UMod operation, using the original dividend and the current divisor.
345 rewriter.replaceOpWithNewOp<spirv::UModOp>(
346 umodOp, umodOp.getType(), prevUMod.getOperand(0), umodOp.getOperand(1));
347
348 return success();
349 }
350};
351
352void spirv::UModOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
353 MLIRContext *context) {
354 patterns.add<UModSimplification>(context);
355}
356
357//===----------------------------------------------------------------------===//
358// spirv.BitcastOp
359//===----------------------------------------------------------------------===//
360
361OpFoldResult spirv::BitcastOp::fold(FoldAdaptor /*adaptor*/) {
362 Value curInput = getOperand();
363 if (getType() == curInput.getType())
364 return curInput;
365
366 // Look through nested bitcasts.
367 if (auto prevCast = curInput.getDefiningOp<spirv::BitcastOp>()) {
368 Value prevInput = prevCast.getOperand();
369 if (prevInput.getType() == getType())
370 return prevInput;
371
372 getOperandMutable().assign(prevInput);
373 return getResult();
374 }
375
376 // TODO(kuhar): Consider constant-folding the operand attribute.
377 return {};
378}
379
380//===----------------------------------------------------------------------===//
381// spirv.CompositeExtractOp
382//===----------------------------------------------------------------------===//
383
384OpFoldResult spirv::CompositeExtractOp::fold(FoldAdaptor adaptor) {
385 Value compositeOp = getComposite();
386
387 while (auto insertOp =
388 compositeOp.getDefiningOp<spirv::CompositeInsertOp>()) {
389 if (getIndices() == insertOp.getIndices())
390 return insertOp.getObject();
391 compositeOp = insertOp.getComposite();
392 }
393
394 if (auto constructOp =
395 compositeOp.getDefiningOp<spirv::CompositeConstructOp>()) {
396 auto type = cast<spirv::CompositeType>(constructOp.getType());
397 if (getIndices().size() == 1 &&
398 constructOp.getConstituents().size() == type.getNumElements()) {
399 auto i = cast<IntegerAttr>(*getIndices().begin());
400 if (i.getValue().getSExtValue() <
401 static_cast<int64_t>(constructOp.getConstituents().size()))
402 return constructOp.getConstituents()[i.getValue().getSExtValue()];
403 }
404 }
405
406 auto indexVector = llvm::map_to_vector(getIndices(), [](Attribute attr) {
407 return static_cast<unsigned>(cast<IntegerAttr>(attr).getInt());
408 });
409 return extractCompositeElement(adaptor.getComposite(), indexVector);
410}
411
412//===----------------------------------------------------------------------===//
413// spirv.Constant
414//===----------------------------------------------------------------------===//
415
416OpFoldResult spirv::ConstantOp::fold(FoldAdaptor /*adaptor*/) {
417 return getValue();
418}
419
420//===----------------------------------------------------------------------===//
421// spirv.IAdd
422//===----------------------------------------------------------------------===//
423
424OpFoldResult spirv::IAddOp::fold(FoldAdaptor adaptor) {
425 // x + 0 = x
426 if (matchPattern(getOperand2(), m_Zero()))
427 return getOperand1();
428
429 // According to the SPIR-V spec:
430 //
431 // The resulting value will equal the low-order N bits of the correct result
432 // R, where N is the component width and R is computed with enough precision
433 // to avoid overflow and underflow.
435 adaptor.getOperands(),
436 [](APInt a, const APInt &b) { return std::move(a) + b; });
437}
438
439//===----------------------------------------------------------------------===//
440// spirv.IMul
441//===----------------------------------------------------------------------===//
442
443OpFoldResult spirv::IMulOp::fold(FoldAdaptor adaptor) {
444 // x * 0 == 0
445 if (matchPattern(getOperand2(), m_Zero()))
446 return getOperand2();
447 // x * 1 = x
448 if (matchPattern(getOperand2(), m_One()))
449 return getOperand1();
450
451 // According to the SPIR-V spec:
452 //
453 // The resulting value will equal the low-order N bits of the correct result
454 // R, where N is the component width and R is computed with enough precision
455 // to avoid overflow and underflow.
457 adaptor.getOperands(),
458 [](const APInt &a, const APInt &b) { return a * b; });
459}
460
461//===----------------------------------------------------------------------===//
462// spirv.ISub
463//===----------------------------------------------------------------------===//
464
465OpFoldResult spirv::ISubOp::fold(FoldAdaptor adaptor) {
466 // x - x = 0
467 if (getOperand1() == getOperand2())
469
470 // According to the SPIR-V spec:
471 //
472 // The resulting value will equal the low-order N bits of the correct result
473 // R, where N is the component width and R is computed with enough precision
474 // to avoid overflow and underflow.
476 adaptor.getOperands(),
477 [](APInt a, const APInt &b) { return std::move(a) - b; });
478}
479
480//===----------------------------------------------------------------------===//
481// spirv.SDiv
482//===----------------------------------------------------------------------===//
483
484OpFoldResult spirv::SDivOp::fold(FoldAdaptor adaptor) {
485 // sdiv (x, 1) = x
486 if (matchPattern(getOperand2(), m_One()))
487 return getOperand1();
488
489 // According to the SPIR-V spec:
490 //
491 // Signed-integer division of Operand 1 divided by Operand 2.
492 // Results are computed per component. Behavior is undefined if Operand 2 is
493 // 0. Behavior is undefined if Operand 2 is -1 and Operand 1 is the minimum
494 // representable value for the operands' type, causing signed overflow.
495 //
496 // So don't fold during undefined behavior.
497 bool div0OrOverflow = false;
499 adaptor.getOperands(), [&](const APInt &a, const APInt &b) {
500 if (div0OrOverflow || isDivZeroOrOverflow(a, b)) {
501 div0OrOverflow = true;
502 return a;
503 }
504 return a.sdiv(b);
505 });
506 return div0OrOverflow ? Attribute() : res;
507}
508
509//===----------------------------------------------------------------------===//
510// spirv.SMod
511//===----------------------------------------------------------------------===//
512
513OpFoldResult spirv::SModOp::fold(FoldAdaptor adaptor) {
514 // smod (x, 1) = 0
515 if (matchPattern(getOperand2(), m_One()))
517
518 // According to SPIR-V spec:
519 //
520 // Signed remainder operation for the remainder whose sign matches the sign
521 // of Operand 2. Behavior is undefined if Operand 2 is 0. Behavior is
522 // undefined if Operand 2 is -1 and Operand 1 is the minimum representable
523 // value for the operands' type, causing signed overflow. Otherwise, the
524 // result is the remainder r of Operand 1 divided by Operand 2 where if
525 // r ≠ 0, the sign of r is the same as the sign of Operand 2.
526 //
527 // So don't fold during undefined behavior
528 bool div0OrOverflow = false;
530 adaptor.getOperands(), [&](const APInt &a, const APInt &b) {
531 if (div0OrOverflow || isDivZeroOrOverflow(a, b)) {
532 div0OrOverflow = true;
533 return a;
534 }
535 APInt c = a.abs().urem(b.abs());
536 if (c.isZero())
537 return c;
538 if (b.isNegative()) {
539 APInt zero = APInt::getZero(c.getBitWidth());
540 return a.isNegative() ? (std::move(zero) - c) : (b + std::move(c));
541 }
542 if (a.isNegative())
543 return b - std::move(c);
544 return c;
545 });
546 return div0OrOverflow ? Attribute() : res;
547}
548
549//===----------------------------------------------------------------------===//
550// spirv.SRem
551//===----------------------------------------------------------------------===//
552
553OpFoldResult spirv::SRemOp::fold(FoldAdaptor adaptor) {
554 // x % 1 = 0
555 if (matchPattern(getOperand2(), m_One()))
557
558 // According to SPIR-V spec:
559 //
560 // Signed remainder operation for the remainder whose sign matches the sign
561 // of Operand 1. Behavior is undefined if Operand 2 is 0. Behavior is
562 // undefined if Operand 2 is -1 and Operand 1 is the minimum representable
563 // value for the operands' type, causing signed overflow. Otherwise, the
564 // result is the remainder r of Operand 1 divided by Operand 2 where if
565 // r ≠ 0, the sign of r is the same as the sign of Operand 1.
566
567 // Don't fold if it would do undefined behavior.
568 bool div0OrOverflow = false;
570 adaptor.getOperands(), [&](APInt a, const APInt &b) {
571 if (div0OrOverflow || isDivZeroOrOverflow(a, b)) {
572 div0OrOverflow = true;
573 return a;
574 }
575 return a.srem(b);
576 });
577 return div0OrOverflow ? Attribute() : res;
578}
579
580//===----------------------------------------------------------------------===//
581// spirv.UDiv
582//===----------------------------------------------------------------------===//
583
584OpFoldResult spirv::UDivOp::fold(FoldAdaptor adaptor) {
585 // udiv (x, 1) = x
586 if (matchPattern(getOperand2(), m_One()))
587 return getOperand1();
588
589 // According to the SPIR-V spec:
590 //
591 // Unsigned-integer division of Operand 1 divided by Operand 2. Behavior is
592 // undefined if Operand 2 is 0.
593 //
594 // So don't fold during undefined behavior.
595 bool div0 = false;
597 adaptor.getOperands(), [&](const APInt &a, const APInt &b) {
598 if (div0 || b.isZero()) {
599 div0 = true;
600 return a;
601 }
602 return a.udiv(b);
603 });
604 return div0 ? Attribute() : res;
605}
606
607//===----------------------------------------------------------------------===//
608// spirv.UMod
609//===----------------------------------------------------------------------===//
610
611OpFoldResult spirv::UModOp::fold(FoldAdaptor adaptor) {
612 // umod (x, 1) = 0
613 if (matchPattern(getOperand2(), m_One()))
615
616 // According to the SPIR-V spec:
617 //
618 // Unsigned modulo operation of Operand 1 modulo Operand 2. Behavior is
619 // undefined if Operand 2 is 0.
620 //
621 // So don't fold during undefined behavior.
622 bool div0 = false;
624 adaptor.getOperands(), [&](const APInt &a, const APInt &b) {
625 if (div0 || b.isZero()) {
626 div0 = true;
627 return a;
628 }
629 return a.urem(b);
630 });
631 return div0 ? Attribute() : res;
632}
633
634//===----------------------------------------------------------------------===//
635// spirv.SNegate
636//===----------------------------------------------------------------------===//
637
638OpFoldResult spirv::SNegateOp::fold(FoldAdaptor adaptor) {
639 // -(-x) = 0 - (0 - x) = x
640 auto op = getOperand();
641 if (auto negateOp = op.getDefiningOp<spirv::SNegateOp>())
642 return negateOp->getOperand(0);
643
644 // According to the SPIR-V spec:
645 //
646 // Signed-integer subtract of Operand from zero.
648 adaptor.getOperands(), [](const APInt &a) {
649 APInt zero = APInt::getZero(a.getBitWidth());
650 return std::move(zero) - a;
651 });
652}
653
654//===----------------------------------------------------------------------===//
655// spirv.NotOp
656//===----------------------------------------------------------------------===//
657
658OpFoldResult spirv::NotOp::fold(spirv::NotOp::FoldAdaptor adaptor) {
659 // !(!x) = x
660 auto op = getOperand();
661 if (auto notOp = op.getDefiningOp<spirv::NotOp>())
662 return notOp->getOperand(0);
663
664 // According to the SPIR-V spec:
665 //
666 // Complement the bits of Operand.
667 return constFoldUnaryOp<IntegerAttr>(adaptor.getOperands(), [&](APInt a) {
668 a.flipAllBits();
669 return a;
670 });
671}
672
673//===----------------------------------------------------------------------===//
674// spirv.LogicalAnd
675//===----------------------------------------------------------------------===//
676
677OpFoldResult spirv::LogicalAndOp::fold(FoldAdaptor adaptor) {
678 if (std::optional<bool> rhs =
679 getScalarOrSplatBoolAttr(adaptor.getOperand2())) {
680 // x && true = x
681 if (*rhs)
682 return getOperand1();
683
684 // x && false = false
685 if (!*rhs)
686 return adaptor.getOperand2();
687 }
688
689 return Attribute();
690}
691
692//===----------------------------------------------------------------------===//
693// spirv.LogicalEqualOp
694//===----------------------------------------------------------------------===//
695
697spirv::LogicalEqualOp::fold(spirv::LogicalEqualOp::FoldAdaptor adaptor) {
698 // x == x -> true
699 if (getOperand1() == getOperand2()) {
700 auto trueAttr = BoolAttr::get(getContext(), true);
701 if (isa<IntegerType>(getType()))
702 return trueAttr;
703 if (auto vecTy = dyn_cast<VectorType>(getType()))
704 return SplatElementsAttr::get(vecTy, trueAttr);
705 }
706
708 adaptor.getOperands(), [](const APInt &a, const APInt &b) {
709 return a == b ? APInt::getAllOnes(1) : APInt::getZero(1);
710 });
711}
712
713//===----------------------------------------------------------------------===//
714// spirv.LogicalNotEqualOp
715//===----------------------------------------------------------------------===//
716
717OpFoldResult spirv::LogicalNotEqualOp::fold(FoldAdaptor adaptor) {
718 if (std::optional<bool> rhs =
719 getScalarOrSplatBoolAttr(adaptor.getOperand2())) {
720 // x != false -> x
721 if (!rhs.value())
722 return getOperand1();
723 }
724
725 // x == x -> false
726 if (getOperand1() == getOperand2()) {
727 auto falseAttr = BoolAttr::get(getContext(), false);
728 if (isa<IntegerType>(getType()))
729 return falseAttr;
730 if (auto vecTy = dyn_cast<VectorType>(getType()))
731 return SplatElementsAttr::get(vecTy, falseAttr);
732 }
733
735 adaptor.getOperands(), [](const APInt &a, const APInt &b) {
736 return a == b ? APInt::getZero(1) : APInt::getAllOnes(1);
737 });
738}
739
740//===----------------------------------------------------------------------===//
741// spirv.LogicalNot
742//===----------------------------------------------------------------------===//
743
744OpFoldResult spirv::LogicalNotOp::fold(FoldAdaptor adaptor) {
745 // !(!x) = x
746 auto op = getOperand();
747 if (auto notOp = op.getDefiningOp<spirv::LogicalNotOp>())
748 return notOp->getOperand(0);
749
750 // According to the SPIR-V spec:
751 //
752 // Complement the bits of Operand.
754 adaptor.getOperands(), [](const APInt &a) {
755 return a == 1 ? APInt::getZero(1) : APInt::getAllOnes(1);
756 });
757}
758
759void spirv::LogicalNotOp::getCanonicalizationPatterns(
760 RewritePatternSet &results, MLIRContext *context) {
761 results
762 .add<ConvertLogicalNotOfIEqual, ConvertLogicalNotOfINotEqual,
763 ConvertLogicalNotOfLogicalEqual, ConvertLogicalNotOfLogicalNotEqual>(
764 context);
765}
766
767//===----------------------------------------------------------------------===//
768// spirv.LogicalOr
769//===----------------------------------------------------------------------===//
770
771OpFoldResult spirv::LogicalOrOp::fold(FoldAdaptor adaptor) {
772 if (auto rhs = getScalarOrSplatBoolAttr(adaptor.getOperand2())) {
773 if (*rhs) {
774 // x || true = true
775 return adaptor.getOperand2();
776 }
777
778 if (!*rhs) {
779 // x || false = x
780 return getOperand1();
781 }
782 }
783
784 return Attribute();
785}
786
787//===----------------------------------------------------------------------===//
788// spirv.SelectOp
789//===----------------------------------------------------------------------===//
790
791OpFoldResult spirv::SelectOp::fold(FoldAdaptor adaptor) {
792 // spirv.Select _ x x -> x
793 Value trueVals = getTrueValue();
794 Value falseVals = getFalseValue();
795 if (trueVals == falseVals)
796 return trueVals;
797
798 ArrayRef<Attribute> operands = adaptor.getOperands();
799
800 // spirv.Select true x y -> x
801 // spirv.Select false x y -> y
802 if (auto boolAttr = getScalarOrSplatBoolAttr(operands[0]))
803 return *boolAttr ? trueVals : falseVals;
804
805 // Check that all the operands are constant
806 if (!operands[0] || !operands[1] || !operands[2])
807 return Attribute();
808
809 // Note: getScalarOrSplatBoolAttr will always return a boolAttr if we are in
810 // the scalar case. Hence, we are only required to consider the case of
811 // DenseElementsAttr in foldSelectOp.
812 auto condAttrs = dyn_cast<DenseElementsAttr>(operands[0]);
813 auto trueAttrs = dyn_cast<DenseElementsAttr>(operands[1]);
814 auto falseAttrs = dyn_cast<DenseElementsAttr>(operands[2]);
815 if (!condAttrs || !trueAttrs || !falseAttrs)
816 return Attribute();
817
818 auto elementResults = llvm::to_vector<4>(trueAttrs.getValues<Attribute>());
819 auto iters = llvm::zip_equal(elementResults, condAttrs.getValues<BoolAttr>(),
820 falseAttrs.getValues<Attribute>());
821 for (auto [result, cond, falseRes] : iters) {
822 if (!cond.getValue())
823 result = falseRes;
824 }
825
826 auto resultType = trueAttrs.getType();
827 return DenseElementsAttr::get(cast<ShapedType>(resultType), elementResults);
828}
829
830//===----------------------------------------------------------------------===//
831// spirv.IEqualOp
832//===----------------------------------------------------------------------===//
833
834OpFoldResult spirv::IEqualOp::fold(spirv::IEqualOp::FoldAdaptor adaptor) {
835 // x == x -> true
836 if (getOperand1() == getOperand2()) {
837 auto trueAttr = BoolAttr::get(getContext(), true);
838 if (isa<IntegerType>(getType()))
839 return trueAttr;
840 if (auto vecTy = dyn_cast<VectorType>(getType()))
841 return SplatElementsAttr::get(vecTy, trueAttr);
842 }
843
845 adaptor.getOperands(), getType(), [](const APInt &a, const APInt &b) {
846 return a == b ? APInt::getAllOnes(1) : APInt::getZero(1);
847 });
848}
849
850//===----------------------------------------------------------------------===//
851// spirv.INotEqualOp
852//===----------------------------------------------------------------------===//
853
854OpFoldResult spirv::INotEqualOp::fold(spirv::INotEqualOp::FoldAdaptor adaptor) {
855 // x == x -> false
856 if (getOperand1() == getOperand2()) {
857 auto falseAttr = BoolAttr::get(getContext(), false);
858 if (isa<IntegerType>(getType()))
859 return falseAttr;
860 if (auto vecTy = dyn_cast<VectorType>(getType()))
861 return SplatElementsAttr::get(vecTy, falseAttr);
862 }
863
865 adaptor.getOperands(), getType(), [](const APInt &a, const APInt &b) {
866 return a == b ? APInt::getZero(1) : APInt::getAllOnes(1);
867 });
868}
869
870//===----------------------------------------------------------------------===//
871// spirv.SGreaterThan
872//===----------------------------------------------------------------------===//
873
875spirv::SGreaterThanOp::fold(spirv::SGreaterThanOp::FoldAdaptor adaptor) {
876 // x == x -> false
877 if (getOperand1() == getOperand2()) {
878 auto falseAttr = BoolAttr::get(getContext(), false);
879 if (isa<IntegerType>(getType()))
880 return falseAttr;
881 if (auto vecTy = dyn_cast<VectorType>(getType()))
882 return SplatElementsAttr::get(vecTy, falseAttr);
883 }
884
886 adaptor.getOperands(), getType(), [](const APInt &a, const APInt &b) {
887 return a.sgt(b) ? APInt::getAllOnes(1) : APInt::getZero(1);
888 });
889}
890
891//===----------------------------------------------------------------------===//
892// spirv.SGreaterThanEqual
893//===----------------------------------------------------------------------===//
894
895OpFoldResult spirv::SGreaterThanEqualOp::fold(
896 spirv::SGreaterThanEqualOp::FoldAdaptor adaptor) {
897 // x == x -> true
898 if (getOperand1() == getOperand2()) {
899 auto trueAttr = BoolAttr::get(getContext(), true);
900 if (isa<IntegerType>(getType()))
901 return trueAttr;
902 if (auto vecTy = dyn_cast<VectorType>(getType()))
903 return SplatElementsAttr::get(vecTy, trueAttr);
904 }
905
907 adaptor.getOperands(), getType(), [](const APInt &a, const APInt &b) {
908 return a.sge(b) ? APInt::getAllOnes(1) : APInt::getZero(1);
909 });
910}
911
912//===----------------------------------------------------------------------===//
913// spirv.UGreaterThan
914//===----------------------------------------------------------------------===//
915
917spirv::UGreaterThanOp::fold(spirv::UGreaterThanOp::FoldAdaptor adaptor) {
918 // x == x -> false
919 if (getOperand1() == getOperand2()) {
920 auto falseAttr = BoolAttr::get(getContext(), false);
921 if (isa<IntegerType>(getType()))
922 return falseAttr;
923 if (auto vecTy = dyn_cast<VectorType>(getType()))
924 return SplatElementsAttr::get(vecTy, falseAttr);
925 }
926
928 adaptor.getOperands(), getType(), [](const APInt &a, const APInt &b) {
929 return a.ugt(b) ? APInt::getAllOnes(1) : APInt::getZero(1);
930 });
931}
932
933//===----------------------------------------------------------------------===//
934// spirv.UGreaterThanEqual
935//===----------------------------------------------------------------------===//
936
937OpFoldResult spirv::UGreaterThanEqualOp::fold(
938 spirv::UGreaterThanEqualOp::FoldAdaptor adaptor) {
939 // x == x -> true
940 if (getOperand1() == getOperand2()) {
941 auto trueAttr = BoolAttr::get(getContext(), true);
942 if (isa<IntegerType>(getType()))
943 return trueAttr;
944 if (auto vecTy = dyn_cast<VectorType>(getType()))
945 return SplatElementsAttr::get(vecTy, trueAttr);
946 }
947
949 adaptor.getOperands(), getType(), [](const APInt &a, const APInt &b) {
950 return a.uge(b) ? APInt::getAllOnes(1) : APInt::getZero(1);
951 });
952}
953
954//===----------------------------------------------------------------------===//
955// spirv.SLessThan
956//===----------------------------------------------------------------------===//
957
958OpFoldResult spirv::SLessThanOp::fold(spirv::SLessThanOp::FoldAdaptor adaptor) {
959 // x == x -> false
960 if (getOperand1() == getOperand2()) {
961 auto falseAttr = BoolAttr::get(getContext(), false);
962 if (isa<IntegerType>(getType()))
963 return falseAttr;
964 if (auto vecTy = dyn_cast<VectorType>(getType()))
965 return SplatElementsAttr::get(vecTy, falseAttr);
966 }
967
969 adaptor.getOperands(), getType(), [](const APInt &a, const APInt &b) {
970 return a.slt(b) ? APInt::getAllOnes(1) : APInt::getZero(1);
971 });
972}
973
974//===----------------------------------------------------------------------===//
975// spirv.SLessThanEqual
976//===----------------------------------------------------------------------===//
977
979spirv::SLessThanEqualOp::fold(spirv::SLessThanEqualOp::FoldAdaptor adaptor) {
980 // x == x -> true
981 if (getOperand1() == getOperand2()) {
982 auto trueAttr = BoolAttr::get(getContext(), true);
983 if (isa<IntegerType>(getType()))
984 return trueAttr;
985 if (auto vecTy = dyn_cast<VectorType>(getType()))
986 return SplatElementsAttr::get(vecTy, trueAttr);
987 }
988
990 adaptor.getOperands(), getType(), [](const APInt &a, const APInt &b) {
991 return a.sle(b) ? APInt::getAllOnes(1) : APInt::getZero(1);
992 });
993}
994
995//===----------------------------------------------------------------------===//
996// spirv.ULessThan
997//===----------------------------------------------------------------------===//
998
999OpFoldResult spirv::ULessThanOp::fold(spirv::ULessThanOp::FoldAdaptor adaptor) {
1000 // x == x -> false
1001 if (getOperand1() == getOperand2()) {
1002 auto falseAttr = BoolAttr::get(getContext(), false);
1003 if (isa<IntegerType>(getType()))
1004 return falseAttr;
1005 if (auto vecTy = dyn_cast<VectorType>(getType()))
1006 return SplatElementsAttr::get(vecTy, falseAttr);
1007 }
1008
1010 adaptor.getOperands(), getType(), [](const APInt &a, const APInt &b) {
1011 return a.ult(b) ? APInt::getAllOnes(1) : APInt::getZero(1);
1012 });
1013}
1014
1015//===----------------------------------------------------------------------===//
1016// spirv.ULessThanEqual
1017//===----------------------------------------------------------------------===//
1018
1020spirv::ULessThanEqualOp::fold(spirv::ULessThanEqualOp::FoldAdaptor adaptor) {
1021 // x == x -> true
1022 if (getOperand1() == getOperand2()) {
1023 auto trueAttr = BoolAttr::get(getContext(), true);
1024 if (isa<IntegerType>(getType()))
1025 return trueAttr;
1026 if (auto vecTy = dyn_cast<VectorType>(getType()))
1027 return SplatElementsAttr::get(vecTy, trueAttr);
1028 }
1029
1031 adaptor.getOperands(), getType(), [](const APInt &a, const APInt &b) {
1032 return a.ule(b) ? APInt::getAllOnes(1) : APInt::getZero(1);
1033 });
1034}
1035
1036//===----------------------------------------------------------------------===//
1037// spirv.ShiftLeftLogical
1038//===----------------------------------------------------------------------===//
1039
1040OpFoldResult spirv::ShiftLeftLogicalOp::fold(
1041 spirv::ShiftLeftLogicalOp::FoldAdaptor adaptor) {
1042 // x << 0 -> x
1043 if (matchPattern(adaptor.getOperand2(), m_Zero())) {
1044 return getOperand1();
1045 }
1046
1047 // Unfortunately due to below undefined behaviour can't fold 0 for Base.
1048
1049 // Results are computed per component, and within each component, per bit...
1050 //
1051 // The result is undefined if Shift is greater than or equal to the bit width
1052 // of the components of Base.
1053 //
1054 // So we can use the APInt << method, but don't fold if undefined behaviour.
1055 bool shiftToLarge = false;
1057 adaptor.getOperands(), [&](const APInt &a, const APInt &b) {
1058 if (shiftToLarge || b.uge(a.getBitWidth())) {
1059 shiftToLarge = true;
1060 return a;
1061 }
1062 return a << b;
1063 });
1064 return shiftToLarge ? Attribute() : res;
1065}
1066
1067//===----------------------------------------------------------------------===//
1068// spirv.ShiftRightArithmetic
1069//===----------------------------------------------------------------------===//
1070
1071OpFoldResult spirv::ShiftRightArithmeticOp::fold(
1072 spirv::ShiftRightArithmeticOp::FoldAdaptor adaptor) {
1073 // x >> 0 -> x
1074 if (matchPattern(adaptor.getOperand2(), m_Zero())) {
1075 return getOperand1();
1076 }
1077
1078 // Unfortunately due to below undefined behaviour can't fold 0, -1 for Base.
1079
1080 // Results are computed per component, and within each component, per bit...
1081 //
1082 // The result is undefined if Shift is greater than or equal to the bit width
1083 // of the components of Base.
1084 //
1085 // So we can use the APInt ashr method, but don't fold if undefined behaviour.
1086 bool shiftToLarge = false;
1088 adaptor.getOperands(), [&](const APInt &a, const APInt &b) {
1089 if (shiftToLarge || b.uge(a.getBitWidth())) {
1090 shiftToLarge = true;
1091 return a;
1092 }
1093 return a.ashr(b);
1094 });
1095 return shiftToLarge ? Attribute() : res;
1096}
1097
1098//===----------------------------------------------------------------------===//
1099// spirv.ShiftRightLogical
1100//===----------------------------------------------------------------------===//
1101
1102OpFoldResult spirv::ShiftRightLogicalOp::fold(
1103 spirv::ShiftRightLogicalOp::FoldAdaptor adaptor) {
1104 // x >> 0 -> x
1105 if (matchPattern(adaptor.getOperand2(), m_Zero())) {
1106 return getOperand1();
1107 }
1108
1109 // Unfortunately due to below undefined behaviour can't fold 0 for Base.
1110
1111 // Results are computed per component, and within each component, per bit...
1112 //
1113 // The result is undefined if Shift is greater than or equal to the bit width
1114 // of the components of Base.
1115 //
1116 // So we can use the APInt lshr method, but don't fold if undefined behaviour.
1117 bool shiftToLarge = false;
1119 adaptor.getOperands(), [&](const APInt &a, const APInt &b) {
1120 if (shiftToLarge || b.uge(a.getBitWidth())) {
1121 shiftToLarge = true;
1122 return a;
1123 }
1124 return a.lshr(b);
1125 });
1126 return shiftToLarge ? Attribute() : res;
1127}
1128
1129//===----------------------------------------------------------------------===//
1130// spirv.BitwiseAndOp
1131//===----------------------------------------------------------------------===//
1132
1134spirv::BitwiseAndOp::fold(spirv::BitwiseAndOp::FoldAdaptor adaptor) {
1135 // x & x -> x
1136 if (getOperand1() == getOperand2()) {
1137 return getOperand1();
1138 }
1139
1140 APInt rhsMask;
1141 if (matchPattern(adaptor.getOperand2(), m_ConstantInt(&rhsMask))) {
1142 // x & 0 -> 0
1143 if (rhsMask.isZero())
1144 return getOperand2();
1145
1146 // x & <all ones> -> x
1147 if (rhsMask.isAllOnes())
1148 return getOperand1();
1149
1150 // (UConvert x : iN to iK) & <mask with N low bits set> -> UConvert x
1151 if (auto zext = getOperand1().getDefiningOp<spirv::UConvertOp>()) {
1152 int valueBits =
1154 if (rhsMask.zextOrTrunc(valueBits).isAllOnes())
1155 return getOperand1();
1156 }
1157 }
1158
1159 // According to the SPIR-V spec:
1160 //
1161 // Type is a scalar or vector of integer type.
1162 // Results are computed per component, and within each component, per bit.
1163 // So we can use the APInt & method.
1165 adaptor.getOperands(),
1166 [](const APInt &a, const APInt &b) { return a & b; });
1167}
1168
1169//===----------------------------------------------------------------------===//
1170// spirv.BitwiseOrOp
1171//===----------------------------------------------------------------------===//
1172
1173OpFoldResult spirv::BitwiseOrOp::fold(spirv::BitwiseOrOp::FoldAdaptor adaptor) {
1174 // x | x -> x
1175 if (getOperand1() == getOperand2()) {
1176 return getOperand1();
1177 }
1178
1179 APInt rhsMask;
1180 if (matchPattern(adaptor.getOperand2(), m_ConstantInt(&rhsMask))) {
1181 // x | 0 -> x
1182 if (rhsMask.isZero())
1183 return getOperand1();
1184
1185 // x | <all ones> -> <all ones>
1186 if (rhsMask.isAllOnes())
1187 return getOperand2();
1188 }
1189
1190 // According to the SPIR-V spec:
1191 //
1192 // Type is a scalar or vector of integer type.
1193 // Results are computed per component, and within each component, per bit.
1194 // So we can use the APInt | method.
1196 adaptor.getOperands(),
1197 [](const APInt &a, const APInt &b) { return a | b; });
1198}
1199
1200//===----------------------------------------------------------------------===//
1201// spirv.BitwiseXorOp
1202//===----------------------------------------------------------------------===//
1203
1205spirv::BitwiseXorOp::fold(spirv::BitwiseXorOp::FoldAdaptor adaptor) {
1206 // x ^ 0 -> x
1207 if (matchPattern(adaptor.getOperand2(), m_Zero())) {
1208 return getOperand1();
1209 }
1210
1211 // x ^ x -> 0
1212 if (getOperand1() == getOperand2())
1214
1215 // According to the SPIR-V spec:
1216 //
1217 // Type is a scalar or vector of integer type.
1218 // Results are computed per component, and within each component, per bit.
1219 // So we can use the APInt ^ method.
1221 adaptor.getOperands(),
1222 [](const APInt &a, const APInt &b) { return a ^ b; });
1223}
1224
1225//===----------------------------------------------------------------------===//
1226// spirv.mlir.selection
1227//===----------------------------------------------------------------------===//
1228
1229namespace {
1230// Blocks from the given `spirv.mlir.selection` operation must satisfy the
1231// following layout:
1232//
1233// +-----------------------------------------------+
1234// | header block |
1235// | spirv.BranchConditionalOp %cond, ^case0, ^case1 |
1236// +-----------------------------------------------+
1237// / \
1238// ...
1239//
1240//
1241// +------------------------+ +------------------------+
1242// | case #0 | | case #1 |
1243// | spirv.Store %ptr %value0 | | spirv.Store %ptr %value1 |
1244// | spirv.Branch ^merge | | spirv.Branch ^merge |
1245// +------------------------+ +------------------------+
1246//
1247//
1248// ...
1249// \ /
1250// v
1251// +-------------+
1252// | merge block |
1253// +-------------+
1254//
1255struct ConvertSelectionOpToSelect final : OpRewritePattern<spirv::SelectionOp> {
1256 using Base::Base;
1257
1258 LogicalResult matchAndRewrite(spirv::SelectionOp selectionOp,
1259 PatternRewriter &rewriter) const override {
1260 Operation *op = selectionOp.getOperation();
1261 Region &body = op->getRegion(0);
1262 // Verifier allows an empty region for `spirv.mlir.selection`.
1263 if (body.empty()) {
1264 return failure();
1265 }
1266
1267 // Check that region consists of 4 blocks:
1268 // header block, `true` block, `false` block and merge block.
1269 if (llvm::range_size(body) != 4) {
1270 return failure();
1271 }
1272
1273 Block *headerBlock = selectionOp.getHeaderBlock();
1274 if (!onlyContainsBranchConditionalOp(headerBlock)) {
1275 return failure();
1276 }
1277
1278 auto brConditionalOp =
1279 cast<spirv::BranchConditionalOp>(headerBlock->front());
1280
1281 Block *trueBlock = brConditionalOp.getSuccessor(0);
1282 Block *falseBlock = brConditionalOp.getSuccessor(1);
1283 Block *mergeBlock = selectionOp.getMergeBlock();
1284
1285 if (failed(canCanonicalizeSelection(trueBlock, falseBlock, mergeBlock)))
1286 return failure();
1287
1288 Value trueValue = getSrcValue(trueBlock);
1289 Value falseValue = getSrcValue(falseBlock);
1290 Value ptrValue = getDstPtr(trueBlock);
1291 auto storeOp = cast<spirv::StoreOp>(trueBlock->front());
1292
1293 auto selectOp = spirv::SelectOp::create(
1294 rewriter, selectionOp.getLoc(), trueValue.getType(),
1295 brConditionalOp.getCondition(), trueValue, falseValue);
1296 auto newStore = spirv::StoreOp::create(
1297 rewriter, selectOp.getLoc(), ptrValue, selectOp.getResult(),
1298 storeOp.getMemoryAccessAttr(), storeOp.getAlignmentAttr());
1299 newStore->setDiscardableAttrs(storeOp->getDiscardableAttrDictionary());
1300
1301 // `spirv.mlir.selection` is not needed anymore.
1302 rewriter.eraseOp(op);
1303 return success();
1304 }
1305
1306private:
1307 // Checks that given blocks follow the following rules:
1308 // 1. Each conditional block consists of two operations, the first operation
1309 // is a `spirv.Store` and the last operation is a `spirv.Branch`.
1310 // 2. Each `spirv.Store` uses the same pointer and the same memory attributes.
1311 // 3. A control flow goes into the given merge block from the given
1312 // conditional blocks.
1313 LogicalResult canCanonicalizeSelection(Block *trueBlock, Block *falseBlock,
1314 Block *mergeBlock) const;
1315
1316 bool onlyContainsBranchConditionalOp(Block *block) const {
1317 return llvm::hasSingleElement(*block) &&
1318 isa<spirv::BranchConditionalOp>(block->front());
1319 }
1320
1321 bool isSameAttrList(spirv::StoreOp lhs, spirv::StoreOp rhs) const {
1322 return lhs->getDiscardableAttrDictionary() ==
1323 rhs->getDiscardableAttrDictionary() &&
1324 lhs.getProperties() == rhs.getProperties();
1325 }
1326
1327 // Returns a source value for the given block.
1328 Value getSrcValue(Block *block) const {
1329 auto storeOp = cast<spirv::StoreOp>(block->front());
1330 return storeOp.getValue();
1331 }
1332
1333 // Returns a destination value for the given block.
1334 Value getDstPtr(Block *block) const {
1335 auto storeOp = cast<spirv::StoreOp>(block->front());
1336 return storeOp.getPtr();
1337 }
1338};
1339
1340LogicalResult ConvertSelectionOpToSelect::canCanonicalizeSelection(
1341 Block *trueBlock, Block *falseBlock, Block *mergeBlock) const {
1342 // Each block must consists of 2 operations.
1343 if (llvm::range_size(*trueBlock) != 2 || llvm::range_size(*falseBlock) != 2) {
1344 return failure();
1345 }
1346
1347 auto trueBrStoreOp = dyn_cast<spirv::StoreOp>(trueBlock->front());
1348 auto trueBrBranchOp =
1349 dyn_cast<spirv::BranchOp>(*std::next(trueBlock->begin()));
1350 auto falseBrStoreOp = dyn_cast<spirv::StoreOp>(falseBlock->front());
1351 auto falseBrBranchOp =
1352 dyn_cast<spirv::BranchOp>(*std::next(falseBlock->begin()));
1353
1354 if (!trueBrStoreOp || !trueBrBranchOp || !falseBrStoreOp ||
1355 !falseBrBranchOp) {
1356 return failure();
1357 }
1358
1359 // Checks that given type is valid for `spirv.SelectOp`.
1360 // According to SPIR-V spec:
1361 // "Before version 1.4, Result Type must be a pointer, scalar, or vector.
1362 // Starting with version 1.4, Result Type can additionally be a composite type
1363 // other than a vector."
1364 bool isScalarOrVector =
1365 cast<spirv::SPIRVType>(trueBrStoreOp.getValue().getType())
1366 .isScalarOrVector();
1367
1368 // Check that each `spirv.Store` uses the same pointer, memory access
1369 // attributes and a valid type of the value.
1370 if ((trueBrStoreOp.getPtr() != falseBrStoreOp.getPtr()) ||
1371 !isSameAttrList(trueBrStoreOp, falseBrStoreOp) || !isScalarOrVector) {
1372 return failure();
1373 }
1374
1375 if ((trueBrBranchOp->getSuccessor(0) != mergeBlock) ||
1376 (falseBrBranchOp->getSuccessor(0) != mergeBlock)) {
1377 return failure();
1378 }
1379
1380 return success();
1381}
1382} // namespace
1383
1384void spirv::SelectionOp::getCanonicalizationPatterns(RewritePatternSet &results,
1385 MLIRContext *context) {
1386 results.add<ConvertSelectionOpToSelect>(context);
1387}
return success()
if(failed(verifyVectorMemoryOp(getOperation(), memrefType, getVectorType()))) return failure()
static Value getZero(OpBuilder &b, Location loc, Type elementType)
Get zero value for an element type.
static uint64_t zext(uint32_t arg)
lhs
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getContext())
ArithmeticExtendedBinaryFold< spirv::ISubBorrowOp > ISubBorrowFold
MulExtendedOpXOne< spirv::SMulExtendedOp > SMulExtendedOpXOne
static Attribute extractCompositeElement(Attribute composite, ArrayRef< unsigned > indices)
MulExtendedFold< spirv::UMulExtendedOp, false > UMulExtendedOpFold
MulExtendedOpXOne< spirv::UMulExtendedOp > UMulExtendedOpXOne
MulExtendedFold< spirv::SMulExtendedOp, true > SMulExtendedOpFold
static std::optional< bool > getScalarOrSplatBoolAttr(Attribute attr)
Returns the boolean value under the hood if the given boolAttr is a scalar or splat vector bool const...
static bool isDivZeroOrOverflow(const APInt &a, const APInt &b)
ArithmeticExtendedBinaryFold< spirv::IAddCarryOp > IAddCarryFold
Attributes are known-constant values of operations.
Definition Attributes.h:25
Block represents an ordered list of Operations.
Definition Block.h:34
Operation & front()
Definition Block.h:178
iterator begin()
Definition Block.h:168
Special case of IntegerAttr to represent boolean integers, i.e., signless i1 integers.
static BoolAttr get(MLIRContext *context, bool value)
This class is a general helper class for creating context-global objects like types,...
Definition Builders.h:51
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
Definition Builders.cpp:275
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class represents a single result from folding an operation.
This provides public APIs that all operations should have.
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
Definition Operation.h:738
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
bool empty()
Definition Region.h:60
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
Operation::operand_range getIndices(Operation *op)
Get the indices that the given load/store operation is operating on.
Definition Utils.cpp:18
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
detail::constant_int_value_binder m_ConstantInt(IntegerAttr::ValueType *bind_value)
Matches a constant holding a scalar/vector/tensor integer (splat) and writes the integer value to bin...
Definition Matchers.h:527
Attribute constFoldBinaryOp(ArrayRef< Attribute > operands, Type resultType, CalculationT &&calculate)
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
detail::constant_int_predicate_matcher m_Zero()
Matches a constant scalar / vector splat / tensor splat integer zero.
Definition Matchers.h:442
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
detail::constant_int_predicate_matcher m_One()
Matches a constant scalar / vector splat / tensor splat integer one.
Definition Matchers.h:478
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
Attribute constFoldUnaryOp(ArrayRef< Attribute > operands, Type resultType, CalculationT &&calculate)
LogicalResult matchAndRewrite(Op op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(MulOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(MulOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(spirv::UModOp umodOp, PatternRewriter &rewriter) const override
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern Base
Type alias to allow derived classes to inherit constructors with using Base::Base;.
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})