13#include "mlir/IR/Diagnostics.h"
14#include "mlir/IR/Matchers.h"
15#include "mlir/IR/PatternMatch.h"
16#include "llvm/ADT/SetVector.h"
17#include "llvm/ADT/SmallBitVector.h"
18#include "llvm/ADT/TypeSwitch.h"
19#include "llvm/Support/KnownBits.h"
24using namespace matchers;
28 return llvm::any_of(op->getOperands(), [op](
auto operand) {
29 return operand.getDefiningOp() == op;
39 ArrayRef<Value> operands, OpBuilder &builder) {
40 OperationState state(loc, name);
41 state.addOperands(operands);
42 state.addTypes(operands[0].getType());
43 return builder.create(state)->getResult(0);
47 return IntegerAttr::get(IntegerType::get(
context, value.getBitWidth()),
53 if (
auto concat = v.getDefiningOp<
ConcatOp>()) {
54 for (
auto op : concat.getOperands())
56 }
else if (
auto repl = v.getDefiningOp<ReplicateOp>()) {
57 for (
size_t i = 0, e = repl.getMultiple(); i != e; ++i)
68 return op->hasAttr(
"sv.attributes");
72template <
typename SubType>
73struct ComplementMatcher {
75 ComplementMatcher(SubType lhs) : lhs(std::move(lhs)) {}
76 bool match(Operation *op) {
77 auto xorOp = dyn_cast<XorOp>(op);
78 return xorOp && xorOp.isBinaryNot() &&
79 mlir::detail::matchOperandOrValueAtIndex(op, 0, lhs);
84template <
typename SubType>
85static inline ComplementMatcher<SubType>
m_Complement(
const SubType &subExpr) {
86 return ComplementMatcher<SubType>(subExpr);
92 assert((isa<AndOp, OrOp, XorOp, AddOp, MulOp>(op) &&
93 "must be commutative operations"));
94 if (op->hasOneUse()) {
95 auto *user = *op->getUsers().begin();
96 return user->getName() == op->getName() &&
97 op->getAttrOfType<UnitAttr>(
"twoState") ==
98 user->getAttrOfType<UnitAttr>(
"twoState") &&
99 op->getBlock() == user->getBlock();
114 auto inputs = op->getOperands();
116 SmallVector<Value, 4> newOperands;
117 SmallVector<Location, 4> newLocations{op->getLoc()};
118 newOperands.reserve(inputs.size());
120 decltype(inputs.begin()) current, end;
123 SmallVector<Element> worklist;
124 worklist.push_back({inputs.begin(), inputs.end()});
125 bool binFlag = op->hasAttrOfType<UnitAttr>(
"twoState");
126 bool changed =
false;
127 while (!worklist.empty()) {
128 auto &element = worklist.back();
131 if (element.current == element.end) {
136 Value value = *element.current++;
137 auto *flattenOp = value.getDefiningOp();
140 if (!flattenOp || flattenOp->getName() != op->getName() ||
141 flattenOp == op || binFlag != op->hasAttrOfType<UnitAttr>(
"twoState") ||
142 flattenOp->getBlock() != op->getBlock()) {
143 newOperands.push_back(value);
148 if (!value.hasOneUse()) {
156 if (flattenOp->getNumOperands() != 2 || !isa<AndOp, OrOp, XorOp>(op) ||
159 newOperands.push_back(value);
167 auto flattenOpInputs = flattenOp->getOperands();
168 worklist.push_back({flattenOpInputs.begin(), flattenOpInputs.end()});
169 newLocations.push_back(flattenOp->getLoc());
175 Value result =
createGenericOp(FusedLoc::get(op->getContext(), newLocations),
176 op->getName(), newOperands, rewriter);
178 result.getDefiningOp()->setAttr(
"twoState", rewriter.getUnitAttr());
186static std::pair<size_t, size_t>
188 size_t originalOpWidth) {
189 auto users = op->getUsers();
191 "getLowestBitAndHighestBitRequired cannot operate on "
192 "a empty list of uses.");
196 size_t lowestBitRequired = narrowTrailingBits ? originalOpWidth - 1 : 0;
197 size_t highestBitRequired = 0;
199 for (
auto *user : users) {
200 if (
auto extractOp = dyn_cast<ExtractOp>(user)) {
201 size_t lowBit = extractOp.getLowBit();
203 cast<IntegerType>(extractOp.getType()).getWidth() + lowBit - 1;
204 highestBitRequired = std::max(highestBitRequired, highBit);
205 lowestBitRequired = std::min(lowestBitRequired, lowBit);
209 highestBitRequired = originalOpWidth - 1;
210 lowestBitRequired = 0;
214 return {lowestBitRequired, highestBitRequired};
219 PatternRewriter &rewriter) {
220 IntegerType opType = dyn_cast<IntegerType>(op.getResult().getType());
226 if (range.second + 1 == opType.getWidth() && range.first == 0)
229 SmallVector<Value> args;
230 auto newType = rewriter.getIntegerType(range.second - range.first + 1);
231 for (
auto inop : op.getOperands()) {
233 if (inop.getType() != op.getType())
234 args.push_back(inop);
236 args.push_back(rewriter.createOrFold<
ExtractOp>(inop.getLoc(), newType,
239 auto newop = OpTy::create(rewriter, op.getLoc(), newType, args);
240 newop->setDialectAttrs(op->getDialectAttrs());
241 if (op.getTwoState())
242 newop.setTwoState(
true);
244 Value newResult = newop.getResult();
246 newResult = rewriter.createOrFold<
ConcatOp>(
247 op.getLoc(), newResult,
249 APInt::getZero(range.first)));
250 if (range.second + 1 < opType.getWidth())
251 newResult = rewriter.createOrFold<
ConcatOp>(
254 rewriter, op.getLoc(),
255 APInt::getZero(opType.getWidth() - range.second - 1)),
257 rewriter.replaceOp(op, newResult);
265OpFoldResult ReplicateOp::fold(FoldAdaptor adaptor) {
270 if (cast<IntegerType>(getType()).
getWidth() ==
271 getInput().getType().getIntOrFloatBitWidth())
275 if (
auto input = dyn_cast_or_null<IntegerAttr>(adaptor.getInput())) {
276 if (input.getValue().getBitWidth() == 1) {
277 if (input.getValue().isZero())
279 APInt::getZero(cast<IntegerType>(getType()).
getWidth()),
282 APInt::getAllOnes(cast<IntegerType>(getType()).
getWidth()),
286 APInt result = APInt::getZeroWidth();
287 for (
auto i = getMultiple(); i != 0; --i)
288 result = result.concat(input.getValue());
295OpFoldResult ParityOp::fold(FoldAdaptor adaptor) {
300 if (
auto input = dyn_cast_or_null<IntegerAttr>(adaptor.getInput()))
301 return getIntAttr(APInt(1, input.getValue().popcount() & 1), getContext());
304 if (hw::getBitWidth(getInput().getType()) == 1)
310LogicalResult ParityOp::canonicalize(
ParityOp op, PatternRewriter &rewriter) {
315 auto isParityZero = [](Value v) {
317 return matchPattern(v, m_ConstantInt(&value)) && value.popcount() % 2 == 0;
322 auto concat = op.getInput().getDefiningOp<
ConcatOp>();
326 auto operands = concat.getInputs();
327 if (operands.size() != 2)
330 if (isParityZero(operands[0])) {
331 replaceOpWithNewOpAndCopyNamehint<ParityOp>(rewriter, op, operands[1],
336 if (isParityZero(operands[1])) {
337 replaceOpWithNewOpAndCopyNamehint<ParityOp>(rewriter, op, operands[0],
352 hw::PEO paramOpcode) {
353 assert(operands.size() == 2 &&
"binary op takes two operands");
354 if (!operands[0] || !operands[1])
359 return hw::ParamExprAttr::get(paramOpcode, cast<TypedAttr>(operands[0]),
360 cast<TypedAttr>(operands[1]));
363OpFoldResult ShlOp::fold(FoldAdaptor adaptor) {
367 if (
auto rhs = dyn_cast_or_null<IntegerAttr>(adaptor.getRhs())) {
368 if (rhs.getValue().isZero())
369 return getOperand(0);
371 unsigned width = getType().getIntOrFloatBitWidth();
372 if (rhs.getValue().uge(width))
373 return getIntAttr(APInt::getZero(width), getContext());
378LogicalResult ShlOp::canonicalize(
ShlOp op, PatternRewriter &rewriter) {
384 if (!matchPattern(op.getRhs(), m_ConstantInt(&value)))
387 unsigned width = cast<IntegerType>(op.getLhs().getType()).getWidth();
388 if (value.ugt(width))
390 unsigned shift = value.getZExtValue();
393 if (width <= shift || shift == 0)
403 replaceOpWithNewOpAndCopyNamehint<ConcatOp>(rewriter, op, extract, zeros);
407OpFoldResult ShrUOp::fold(FoldAdaptor adaptor) {
411 if (
auto rhs = dyn_cast_or_null<IntegerAttr>(adaptor.getRhs())) {
412 if (rhs.getValue().isZero())
413 return getOperand(0);
415 unsigned width = getType().getIntOrFloatBitWidth();
416 if (rhs.getValue().uge(width))
417 return getIntAttr(APInt::getZero(width), getContext());
422LogicalResult ShrUOp::canonicalize(
ShrUOp op, PatternRewriter &rewriter) {
428 if (!matchPattern(op.getRhs(), m_ConstantInt(&value)))
431 unsigned width = cast<IntegerType>(op.getLhs().getType()).getWidth();
432 if (value.ugt(width))
434 unsigned shift = value.getZExtValue();
437 if (width <= shift || shift == 0)
447 replaceOpWithNewOpAndCopyNamehint<ConcatOp>(rewriter, op, zeros, extract);
451OpFoldResult ShrSOp::fold(FoldAdaptor adaptor) {
455 if (
auto rhs = dyn_cast_or_null<IntegerAttr>(adaptor.getRhs()))
456 if (rhs.getValue().isZero())
457 return getOperand(0);
461LogicalResult ShrSOp::canonicalize(
ShrSOp op, PatternRewriter &rewriter) {
467 if (!matchPattern(op.getRhs(), m_ConstantInt(&value)))
470 unsigned width = cast<IntegerType>(op.getLhs().getType()).getWidth();
471 if (value.ugt(width))
473 unsigned shift = value.getZExtValue();
476 rewriter.createOrFold<
ExtractOp>(op.getLoc(), op.getLhs(), width - 1, 1);
477 auto sext = rewriter.createOrFold<ReplicateOp>(op.getLoc(), topbit, shift);
479 if (width == shift) {
487 replaceOpWithNewOpAndCopyNamehint<ConcatOp>(rewriter, op, sext, extract);
495OpFoldResult ExtractOp::fold(FoldAdaptor adaptor) {
500 if (getInput().getType() == getType())
504 if (
auto input = dyn_cast_or_null<IntegerAttr>(adaptor.getInput())) {
505 unsigned dstWidth = cast<IntegerType>(getType()).getWidth();
506 return getIntAttr(input.getValue().lshr(getLowBit()).trunc(dstWidth),
522 PatternRewriter &rewriter,
523 ArrayRef<size_t> prefixWidths = {}) {
524 auto concatInputs = innerCat.getInputs();
525 size_t numOperands = concatInputs.size();
526 size_t lowBit = op.getLowBit();
531 if (!prefixWidths.empty()) {
534 std::upper_bound(prefixWidths.begin(), prefixWidths.end(), lowBit);
535 assert(it != prefixWidths.end());
536 firstIdx = it - prefixWidths.begin();
537 beginOfFirst = (firstIdx > 0) ? prefixWidths[firstIdx - 1] : 0;
542 for (
size_t i = 0; i < numOperands; ++i) {
544 concatInputs[numOperands - 1 - i].getType().getIntOrFloatBitWidth();
545 if (lowBit < beginOfFirst + w) {
553 SmallVector<Value> reverseConcatArgs;
554 size_t widthRemaining = op.getType().getIntOrFloatBitWidth();
555 size_t extractLo = lowBit - beginOfFirst;
560 for (
size_t i = firstIdx; widthRemaining != 0 && i < numOperands; ++i) {
561 Value concatArg = concatInputs[numOperands - 1 - i];
562 size_t operandWidth = concatArg.getType().getIntOrFloatBitWidth();
563 size_t widthToConsume = std::min(widthRemaining, operandWidth - extractLo);
565 if (widthToConsume == operandWidth && extractLo == 0) {
566 reverseConcatArgs.push_back(concatArg);
568 auto resultType = IntegerType::get(rewriter.getContext(), widthToConsume);
570 rewriter, op.getLoc(), resultType, concatArg, extractLo));
573 widthRemaining -= widthToConsume;
578 if (reverseConcatArgs.size() == 1) {
581 replaceOpWithNewOpAndCopyNamehint<ConcatOp>(
582 rewriter, op, SmallVector<Value>(llvm::reverse(reverseConcatArgs)));
589 PatternRewriter &rewriter) {
590 auto extractResultWidth = cast<IntegerType>(op.getType()).getWidth();
591 auto replicateEltWidth =
592 replicate.getOperand().getType().getIntOrFloatBitWidth();
596 if (op.getLowBit() % replicateEltWidth == 0 &&
597 extractResultWidth % replicateEltWidth == 0) {
598 replaceOpWithNewOpAndCopyNamehint<ReplicateOp>(rewriter, op, op.getType(),
599 replicate.getOperand());
605 if (op.getLowBit() % replicateEltWidth + extractResultWidth <=
607 replaceOpWithNewOpAndCopyNamehint<ExtractOp>(
608 rewriter, op, op.getType(), replicate.getOperand(),
609 op.getLowBit() % replicateEltWidth);
618LogicalResult ExtractOp::canonicalize(
ExtractOp op, PatternRewriter &rewriter) {
621 auto *inputOp = op.getInput().getDefiningOp();
628 .extractBits(cast<IntegerType>(op.getType()).getWidth(),
630 if (knownBits.isConstant()) {
631 replaceOpWithNewOpAndCopyNamehint<hw::ConstantOp>(rewriter, op,
632 knownBits.getConstant());
638 if (
auto innerExtract = dyn_cast_or_null<ExtractOp>(inputOp)) {
639 replaceOpWithNewOpAndCopyNamehint<ExtractOp>(
640 rewriter, op, op.getType(), innerExtract.getInput(),
641 innerExtract.getLowBit() + op.getLowBit());
646 if (
auto innerCat = dyn_cast_or_null<ConcatOp>(inputOp))
650 if (
auto replicate = dyn_cast_or_null<ReplicateOp>(inputOp))
656 if (inputOp && inputOp->getNumOperands() == 2 &&
657 isa<AndOp, OrOp, XorOp>(inputOp)) {
658 if (
auto cstRHS = inputOp->getOperand(1).getDefiningOp<
hw::ConstantOp>()) {
659 auto extractedCst = cstRHS.getValue().extractBits(
660 cast<IntegerType>(op.getType()).getWidth(), op.getLowBit());
661 if (isa<OrOp, XorOp>(inputOp) && extractedCst.isZero()) {
662 replaceOpWithNewOpAndCopyNamehint<ExtractOp>(
663 rewriter, op, op.getType(), inputOp->getOperand(0), op.getLowBit());
671 if (isa<AndOp>(inputOp)) {
674 unsigned lz = extractedCst.countLeadingZeros();
675 unsigned tz = extractedCst.countTrailingZeros();
676 unsigned pop = extractedCst.popcount();
677 if (extractedCst.getBitWidth() - lz - tz == pop) {
678 auto resultTy = rewriter.getIntegerType(pop);
679 SmallVector<Value> resultElts;
682 APInt::getZero(lz)));
683 resultElts.push_back(rewriter.createOrFold<
ExtractOp>(
684 op.getLoc(), resultTy, inputOp->getOperand(0),
685 op.getLowBit() + tz));
688 APInt::getZero(tz)));
689 replaceOpWithNewOpAndCopyNamehint<ConcatOp>(rewriter, op, resultElts);
698 if (cast<IntegerType>(op.getType()).getWidth() == 1 && inputOp)
699 if (
auto shlOp = dyn_cast<ShlOp>(inputOp)) {
701 if (shlOp->hasOneUse())
703 if (lhsCst.getValue().isOne()) {
705 rewriter, shlOp.getLoc(),
706 APInt(lhsCst.getValue().getBitWidth(), op.getLowBit()));
707 replaceOpWithNewOpAndCopyNamehint<ICmpOp>(
708 rewriter, op, ICmpPredicate::eq, shlOp->getOperand(1), newCst,
724 hw::PEO paramOpcode) {
725 assert(operands.size() > 1 &&
"caller should handle one-operand case");
728 if (!operands[1] || !operands[0])
732 if (llvm::all_of(operands.drop_front(2),
733 [&](Attribute in) { return !!in; })) {
734 SmallVector<mlir::TypedAttr> typedOperands;
735 typedOperands.reserve(operands.size());
736 for (
auto operand : operands) {
737 if (
auto typedOperand = dyn_cast<mlir::TypedAttr>(operand))
738 typedOperands.push_back(typedOperand);
742 if (typedOperands.size() == operands.size())
743 return hw::ParamExprAttr::get(paramOpcode, typedOperands);
759 size_t concatIdx,
const APInt &cst,
760 PatternRewriter &rewriter) {
761 auto concatOp = logicalOp->getOperand(concatIdx).getDefiningOp<
ConcatOp>();
762 assert((isa<AndOp, OrOp, XorOp>(logicalOp) && concatOp));
767 llvm::any_of(concatOp->getOperands(), [&](Value operand) ->
bool {
768 auto *operandOp = operand.getDefiningOp();
773 if (isa<hw::ConstantOp>(operandOp))
777 return operandOp->getName() == logicalOp->getName() &&
778 operandOp->hasOneUse() && operandOp->getNumOperands() != 0 &&
779 operandOp->getOperands().back().getDefiningOp<hw::ConstantOp>();
787 auto createLogicalOp = [&](ArrayRef<Value> operands) -> Value {
788 return createGenericOp(logicalOp->getLoc(), logicalOp->getName(), operands,
795 SmallVector<Value> newConcatOperands;
796 newConcatOperands.reserve(concatOp->getNumOperands());
799 size_t nextOperandBit = concatOp.getType().getIntOrFloatBitWidth();
800 for (Value operand : concatOp->getOperands()) {
801 size_t operandWidth = operand.getType().getIntOrFloatBitWidth();
802 nextOperandBit -= operandWidth;
806 cst.lshr(nextOperandBit).trunc(operandWidth));
808 newConcatOperands.push_back(createLogicalOp({operand, eltCst}));
813 ConcatOp::create(rewriter, concatOp.getLoc(), newConcatOperands);
817 if (logicalOp->getNumOperands() > 2) {
818 auto origOperands = logicalOp->getOperands();
819 SmallVector<Value> operands;
821 operands.append(origOperands.begin(), origOperands.begin() + concatIdx);
823 operands.append(origOperands.begin() + concatIdx + 1,
824 origOperands.begin() + (origOperands.size() - 1));
826 operands.push_back(newResult);
827 newResult = createLogicalOp(operands);
837 llvm::SmallDenseSet<std::tuple<ICmpPredicate, Value, Value>> seenPredicates;
839 for (
auto op : operands) {
840 if (
auto icmpOp = op.getDefiningOp<ICmpOp>();
841 icmpOp && icmpOp.getTwoState()) {
842 auto predicate = icmpOp.getPredicate();
843 auto lhs = icmpOp.getLhs();
844 auto rhs = icmpOp.getRhs();
845 if (seenPredicates.contains(
846 {ICmpOp::getNegatedPredicate(predicate), lhs, rhs}))
849 seenPredicates.insert({predicate, lhs, rhs});
855OpFoldResult AndOp::fold(FoldAdaptor adaptor) {
859 APInt value = APInt::getAllOnes(cast<IntegerType>(getType()).
getWidth());
861 auto inputs = adaptor.getInputs();
864 for (
auto operand : inputs) {
865 auto attr = dyn_cast_or_null<IntegerAttr>(operand);
868 value &= attr.getValue();
874 if (inputs.size() == 2)
875 if (
auto intAttr = dyn_cast_or_null<IntegerAttr>(inputs[1]))
876 if (intAttr.getValue().isAllOnes())
877 return getInputs()[0];
880 if (llvm::all_of(getInputs(),
881 [&](
auto in) {
return in == this->getInputs()[0]; }))
882 return getInputs()[0];
885 for (Value arg : getInputs()) {
888 for (Value arg2 : getInputs())
891 APInt::getZero(cast<IntegerType>(getType()).
getWidth()),
912template <
typename Op>
914 if (!op.getType().isInteger(1))
917 auto inputs = op.getInputs();
918 size_t size = inputs.size();
920 auto sourceOp = inputs[0].template getDefiningOp<ExtractOp>();
923 Value source = sourceOp.getOperand();
926 if (size != source.getType().getIntOrFloatBitWidth())
930 llvm::BitVector bits(size);
931 bits.set(sourceOp.getLowBit());
933 for (
size_t i = 1; i != size; ++i) {
934 auto extractOp = inputs[i].template getDefiningOp<ExtractOp>();
935 if (!extractOp || extractOp.getOperand() != source)
937 bits.set(extractOp.getLowBit());
940 return bits.all() ? source : Value();
947template <
typename Op>
950 constexpr unsigned limit = 3;
951 auto inputs = op.getInputs();
954 llvm::SmallDenseSet<Op, 8> checked;
961 llvm::SmallVector<OpWithDepth, 8> worklist;
963 auto enqueue = [&worklist, &checked, &op](Value input,
unsigned depth) {
967 if (depth < limit && input.getParentBlock() == op->getBlock()) {
968 auto inputOp = input.template getDefiningOp<Op>();
969 if (inputOp && inputOp.getTwoState() == op.getTwoState() &&
970 checked.insert(inputOp).second)
971 worklist.push_back({inputOp, depth + 1});
975 for (
auto input : uniqueInputs)
978 while (!worklist.empty()) {
979 auto item = worklist.pop_back_val();
981 for (
auto input : item.op.getInputs()) {
982 uniqueInputs.remove(input);
983 enqueue(input, item.depth);
987 if (uniqueInputs.size() < inputs.size()) {
988 replaceOpWithNewOpAndCopyNamehint<Op>(rewriter, op, op.getType(),
989 uniqueInputs.getArrayRef(),
997LogicalResult AndOp::canonicalize(
AndOp op, PatternRewriter &rewriter) {
1001 auto inputs = op.getInputs();
1002 auto size = inputs.size();
1014 assert(size > 1 &&
"expected 2 or more operands, `fold` should handle this");
1018 if (matchPattern(inputs.back(), m_ConstantInt(&value))) {
1020 if (value.isAllOnes()) {
1021 replaceOpWithNewOpAndCopyNamehint<AndOp>(rewriter, op, op.getType(),
1022 inputs.drop_back(),
false);
1030 if (matchPattern(inputs[size - 2], m_ConstantInt(&value2))) {
1032 SmallVector<Value, 4> newOperands(inputs.drop_back(2));
1033 newOperands.push_back(cst);
1034 replaceOpWithNewOpAndCopyNamehint<AndOp>(rewriter, op, op.getType(),
1035 newOperands,
false);
1040 if (size == 2 && value.isPowerOf2()) {
1045 if (
auto replicate = inputs[0].getDefiningOp<ReplicateOp>()) {
1046 auto replicateOperand = replicate.getOperand();
1047 if (replicateOperand.getType().isInteger(1)) {
1048 unsigned resultWidth = op.getType().getIntOrFloatBitWidth();
1049 auto trailingZeros = value.countTrailingZeros();
1052 SmallVector<Value, 3> concatOperands;
1053 if (trailingZeros != resultWidth - 1) {
1055 rewriter, op.getLoc(),
1056 APInt::getZero(resultWidth - trailingZeros - 1));
1057 concatOperands.push_back(highZeros);
1059 concatOperands.push_back(replicateOperand);
1060 if (trailingZeros != 0) {
1062 rewriter, op.getLoc(), APInt::getZero(trailingZeros));
1063 concatOperands.push_back(lowZeros);
1065 replaceOpWithNewOpAndCopyNamehint<ConcatOp>(
1066 rewriter, op, op.getType(), concatOperands);
1075 unsigned leadingZeros = value.countLeadingZeros();
1076 unsigned trailingZeros = value.countTrailingZeros();
1077 if (leadingZeros > 0 || trailingZeros > 0) {
1078 unsigned maskLength = value.getBitWidth() - leadingZeros - trailingZeros;
1081 SmallVector<Value> operands;
1082 for (
auto input : inputs.drop_back()) {
1083 unsigned offset = trailingZeros;
1084 while (
auto extractOp = input.getDefiningOp<
ExtractOp>()) {
1085 input = extractOp.getInput();
1086 offset += extractOp.getLowBit();
1089 offset, maskLength));
1093 auto narrowMask = value.extractBits(maskLength, trailingZeros);
1094 if (!narrowMask.isAllOnes())
1096 rewriter, inputs.back().getLoc(), narrowMask));
1099 Value narrowValue = operands.back();
1100 if (operands.size() > 1)
1102 AndOp::create(rewriter, op.getLoc(), operands, op.getTwoState());
1106 if (leadingZeros > 0)
1108 rewriter, op.getLoc(), APInt::getZero(leadingZeros)));
1109 operands.push_back(narrowValue);
1110 if (trailingZeros > 0)
1112 rewriter, op.getLoc(), APInt::getZero(trailingZeros)));
1113 replaceOpWithNewOpAndCopyNamehint<ConcatOp>(rewriter, op, operands);
1120 for (
size_t i = 0; i < size - 1; ++i) {
1121 if (
auto concat = inputs[i].getDefiningOp<ConcatOp>())
1135 replaceOpWithNewOpAndCopyNamehint<ICmpOp>(rewriter, op, ICmpPredicate::eq,
1136 source, cmpAgainst);
1141 if (op.getTwoState() && op.getNumOperands() == 2) {
1142 auto isReplicateOfI1 = [](Value v) {
1143 auto rep = v.getDefiningOp<ReplicateOp>();
1146 return rep.getOperand().getType().isInteger(1);
1148 Value x = op.getOperand(0);
1149 Value y = op.getOperand(1);
1150 if (isReplicateOfI1(x))
1152 if (isReplicateOfI1(y)) {
1153 Value p = y.getDefiningOp<ReplicateOp>().getInput();
1155 rewriter, op.getLoc(), rewriter.getIntegerAttr(op.getType(), 0));
1156 replaceOpWithNewOpAndCopyNamehint<MuxOp>(rewriter, op, p, x, zero,
1166OpFoldResult OrOp::fold(FoldAdaptor adaptor) {
1170 auto value = APInt::getZero(cast<IntegerType>(getType()).
getWidth());
1171 auto inputs = adaptor.getInputs();
1173 for (
auto operand : inputs) {
1174 auto attr = dyn_cast_or_null<IntegerAttr>(operand);
1177 value |= attr.getValue();
1178 if (value.isAllOnes())
1183 if (inputs.size() == 2)
1184 if (
auto intAttr = dyn_cast_or_null<IntegerAttr>(inputs[1]))
1185 if (intAttr.getValue().isZero())
1186 return getInputs()[0];
1189 if (llvm::all_of(getInputs(),
1190 [&](
auto in) {
return in == this->getInputs()[0]; }))
1191 return getInputs()[0];
1194 for (Value arg : getInputs()) {
1196 if (matchPattern(arg,
m_Complement(m_Any(&subExpr)))) {
1197 for (Value arg2 : getInputs())
1198 if (arg2 == subExpr)
1200 APInt::getAllOnes(cast<IntegerType>(getType()).
getWidth()),
1210 APInt::getAllOnes(cast<IntegerType>(getType()).
getWidth()),
1217LogicalResult OrOp::canonicalize(
OrOp op, PatternRewriter &rewriter) {
1221 auto inputs = op.getInputs();
1222 auto size = inputs.size();
1234 assert(size > 1 &&
"expected 2 or more operands");
1238 if (matchPattern(inputs.back(), m_ConstantInt(&value))) {
1240 if (value.isZero()) {
1241 replaceOpWithNewOpAndCopyNamehint<OrOp>(rewriter, op, op.getType(),
1242 inputs.drop_back());
1248 if (matchPattern(inputs[size - 2], m_ConstantInt(&value2))) {
1250 SmallVector<Value, 4> newOperands(inputs.drop_back(2));
1251 newOperands.push_back(cst);
1252 replaceOpWithNewOpAndCopyNamehint<OrOp>(rewriter, op, op.getType(),
1260 for (
size_t i = 0; i < size - 1; ++i) {
1261 if (
auto concat = inputs[i].getDefiningOp<ConcatOp>())
1275 replaceOpWithNewOpAndCopyNamehint<ICmpOp>(rewriter, op, ICmpPredicate::ne,
1276 source, cmpAgainst);
1282 if (
auto firstMux = op.getOperand(0).getDefiningOp<
comb::MuxOp>()) {
1284 if (op.getTwoState() && firstMux.getTwoState() &&
1285 matchPattern(firstMux.getFalseValue(), m_ConstantInt(&value)) &&
1287 SmallVector<Value> conditions{firstMux.getCond()};
1288 auto check = [&](Value v) {
1292 conditions.push_back(mux.getCond());
1293 return mux.getTwoState() &&
1294 firstMux.getTrueValue() == mux.getTrueValue() &&
1295 firstMux.getFalseValue() == mux.getFalseValue();
1297 if (llvm::all_of(op.getOperands().drop_front(), check)) {
1298 auto cond = comb::OrOp::create(rewriter, op.getLoc(), conditions,
true);
1299 replaceOpWithNewOpAndCopyNamehint<comb::MuxOp>(
1300 rewriter, op, cond, firstMux.getTrueValue(),
1301 firstMux.getFalseValue(),
true);
1311OpFoldResult XorOp::fold(FoldAdaptor adaptor) {
1315 auto size = getInputs().size();
1316 auto inputs = adaptor.getInputs();
1320 return getInputs()[0];
1323 if (size == 2 && getInputs()[0] == getInputs()[1])
1324 return IntegerAttr::get(getType(), 0);
1327 if (inputs.size() == 2)
1328 if (
auto intAttr = dyn_cast_or_null<IntegerAttr>(inputs[1]))
1329 if (intAttr.getValue().isZero())
1330 return getInputs()[0];
1336 subExpr != getResult())
1345 PatternRewriter &rewriter) {
1346 auto icmp = op.getOperand(icmpOperand).getDefiningOp<ICmpOp>();
1347 auto negatedPred = ICmpOp::getNegatedPredicate(icmp.getPredicate());
1350 ICmpOp::create(rewriter, icmp.getLoc(), negatedPred, icmp.getOperand(0),
1351 icmp.getOperand(1), icmp.getTwoState());
1354 if (op.getNumOperands() > 2) {
1355 SmallVector<Value, 4> newOperands(op.getOperands());
1356 newOperands.pop_back();
1357 newOperands.erase(newOperands.begin() + icmpOperand);
1358 newOperands.push_back(result);
1360 XorOp::create(rewriter, op.getLoc(), newOperands, op.getTwoState());
1366LogicalResult XorOp::canonicalize(
XorOp op, PatternRewriter &rewriter) {
1370 auto inputs = op.getInputs();
1371 auto size = inputs.size();
1372 assert(size > 1 &&
"expected 2 or more operands");
1375 if (inputs[size - 1] == inputs[size - 2]) {
1377 "expected idempotent case for 2 elements handled already.");
1378 replaceOpWithNewOpAndCopyNamehint<XorOp>(rewriter, op, op.getType(),
1379 inputs.drop_back(2),
false);
1385 if (matchPattern(inputs.back(), m_ConstantInt(&value))) {
1387 if (value.isZero()) {
1388 replaceOpWithNewOpAndCopyNamehint<XorOp>(rewriter, op, op.getType(),
1389 inputs.drop_back(),
false);
1395 if (matchPattern(inputs[size - 2], m_ConstantInt(&value2))) {
1397 SmallVector<Value, 4> newOperands(inputs.drop_back(2));
1398 newOperands.push_back(cst);
1399 replaceOpWithNewOpAndCopyNamehint<XorOp>(rewriter, op, op.getType(),
1400 newOperands,
false);
1404 bool isSingleBit = value.getBitWidth() == 1;
1407 for (
size_t i = 0; i < size - 1; ++i) {
1408 Value operand = inputs[i];
1414 if (
auto concat = operand.getDefiningOp<
ConcatOp>())
1419 if (isSingleBit && operand.hasOneUse()) {
1420 assert(value == 1 &&
"single bit constant has to be one if not zero");
1421 if (
auto icmp = operand.getDefiningOp<ICmpOp>())
1429 Value complementVal;
1432 if (matchPattern(op.getResult(),
m_Complement(m_Any(&complementVal))) &&
1433 matchPattern(complementVal, m_SextBy(m_Any(&signExtBits)))) {
1435 auto baseWidth = op.getType().getIntOrFloatBitWidth() -
1436 signExtBits.getType().getIntOrFloatBitWidth();
1458 replaceOpWithNewOpAndCopyNamehint<ParityOp>(rewriter, op, source);
1465OpFoldResult SubOp::fold(FoldAdaptor adaptor) {
1470 if (getRhs() == getLhs())
1472 APInt::getZero(getLhs().getType().getIntOrFloatBitWidth()),
1475 if (adaptor.getRhs()) {
1477 if (adaptor.getLhs()) {
1480 APInt::getAllOnes(getLhs().getType().getIntOrFloatBitWidth()),
1482 auto rhsNeg = hw::ParamExprAttr::get(
1483 hw::PEO::Mul, cast<TypedAttr>(adaptor.getRhs()), negOne);
1484 return hw::ParamExprAttr::get(hw::PEO::Add,
1485 cast<TypedAttr>(adaptor.getLhs()), rhsNeg);
1489 if (
auto rhsC = dyn_cast<IntegerAttr>(adaptor.getRhs())) {
1490 if (rhsC.getValue().isZero())
1498LogicalResult SubOp::canonicalize(
SubOp op, PatternRewriter &rewriter) {
1504 if (matchPattern(op.getRhs(), m_ConstantInt(&value))) {
1506 replaceOpWithNewOpAndCopyNamehint<AddOp>(rewriter, op, op.getLhs(), negCst,
1518OpFoldResult AddOp::fold(FoldAdaptor adaptor) {
1522 auto size = getInputs().size();
1526 return getInputs()[0];
1532LogicalResult AddOp::canonicalize(
AddOp op, PatternRewriter &rewriter) {
1536 auto inputs = op.getInputs();
1537 auto size = inputs.size();
1538 assert(size > 1 &&
"expected 2 or more operands");
1540 APInt value, value2;
1543 if (matchPattern(inputs.back(), m_ConstantInt(&value)) && value.isZero()) {
1544 replaceOpWithNewOpAndCopyNamehint<AddOp>(rewriter, op, op.getType(),
1545 inputs.drop_back(),
false);
1550 if (matchPattern(inputs[size - 1], m_ConstantInt(&value)) &&
1551 matchPattern(inputs[size - 2], m_ConstantInt(&value2))) {
1553 SmallVector<Value, 4> newOperands(inputs.drop_back(2));
1554 newOperands.push_back(cst);
1555 replaceOpWithNewOpAndCopyNamehint<AddOp>(rewriter, op, op.getType(),
1556 newOperands,
false);
1561 if (inputs[size - 1] == inputs[size - 2]) {
1562 SmallVector<Value, 4> newOperands(inputs.drop_back(2));
1566 comb::ShlOp::create(rewriter, op.getLoc(), inputs.back(), one,
false);
1568 newOperands.push_back(shiftLeftOp);
1569 replaceOpWithNewOpAndCopyNamehint<AddOp>(rewriter, op, op.getType(),
1570 newOperands,
false);
1574 auto shlOp = inputs[size - 1].getDefiningOp<
comb::ShlOp>();
1576 if (shlOp && shlOp.getLhs() == inputs[size - 2] &&
1577 matchPattern(shlOp.getRhs(), m_ConstantInt(&value))) {
1579 APInt one(value.getBitWidth(), 1,
false);
1583 std::array<Value, 2> factors = {shlOp.getLhs(), rhs};
1584 auto mulOp = comb::MulOp::create(rewriter, op.getLoc(), factors,
false);
1586 SmallVector<Value, 4> newOperands(inputs.drop_back(2));
1587 newOperands.push_back(mulOp);
1588 replaceOpWithNewOpAndCopyNamehint<AddOp>(rewriter, op, op.getType(),
1589 newOperands,
false);
1593 auto mulOp = inputs[size - 1].getDefiningOp<
comb::MulOp>();
1595 if (mulOp && mulOp.getInputs().size() == 2 &&
1596 mulOp.getInputs()[0] == inputs[size - 2] &&
1597 matchPattern(mulOp.getInputs()[1], m_ConstantInt(&value))) {
1599 APInt one(value.getBitWidth(), 1,
false);
1601 std::array<Value, 2> factors = {mulOp.getInputs()[0], rhs};
1602 auto newMulOp = comb::MulOp::create(rewriter, op.getLoc(), factors,
false);
1604 SmallVector<Value, 4> newOperands(inputs.drop_back(2));
1605 newOperands.push_back(newMulOp);
1606 replaceOpWithNewOpAndCopyNamehint<AddOp>(rewriter, op, op.getType(),
1607 newOperands,
false);
1620 auto addOp = inputs[0].getDefiningOp<
comb::AddOp>();
1621 if (addOp && addOp.getInputs().size() == 2 &&
1622 matchPattern(addOp.getInputs()[1], m_ConstantInt(&value2)) &&
1623 inputs.size() == 2 && matchPattern(inputs[1], m_ConstantInt(&value))) {
1626 replaceOpWithNewOpAndCopyNamehint<AddOp>(
1627 rewriter, op, op.getType(), ArrayRef<Value>{addOp.getInputs()[0], rhs},
1628 op.getTwoState() && addOp.getTwoState());
1635OpFoldResult MulOp::fold(FoldAdaptor adaptor) {
1639 auto size = getInputs().size();
1640 auto inputs = adaptor.getInputs();
1644 return getInputs()[0];
1646 auto width = cast<IntegerType>(getType()).getWidth();
1648 return getIntAttr(APInt::getZero(0), getContext());
1650 APInt value(width, 1,
false);
1653 for (
auto operand : inputs) {
1654 auto attr = dyn_cast_or_null<IntegerAttr>(operand);
1657 value *= attr.getValue();
1666LogicalResult MulOp::canonicalize(
MulOp op, PatternRewriter &rewriter) {
1670 auto inputs = op.getInputs();
1671 auto size = inputs.size();
1672 assert(size > 1 &&
"expected 2 or more operands");
1674 APInt value, value2;
1677 if (size == 2 && matchPattern(inputs.back(), m_ConstantInt(&value)) &&
1678 value.isPowerOf2()) {
1680 value.exactLogBase2());
1682 comb::ShlOp::create(rewriter, op.getLoc(), inputs[0], shift,
false);
1684 replaceOpWithNewOpAndCopyNamehint<MulOp>(rewriter, op, op.getType(),
1685 ArrayRef<Value>(shlOp),
false);
1690 if (matchPattern(inputs.back(), m_ConstantInt(&value)) && value.isOne()) {
1691 replaceOpWithNewOpAndCopyNamehint<MulOp>(rewriter, op, op.getType(),
1692 inputs.drop_back());
1697 if (matchPattern(inputs[size - 1], m_ConstantInt(&value)) &&
1698 matchPattern(inputs[size - 2], m_ConstantInt(&value2))) {
1700 SmallVector<Value, 4> newOperands(inputs.drop_back(2));
1701 newOperands.push_back(cst);
1702 replaceOpWithNewOpAndCopyNamehint<MulOp>(rewriter, op, op.getType(),
1718template <
class Op,
bool isSigned>
1719static OpFoldResult
foldDiv(Op op, ArrayRef<Attribute> constants) {
1720 if (
auto rhsValue = dyn_cast_or_null<IntegerAttr>(constants[1])) {
1722 if (rhsValue.getValue() == 1)
1726 if (rhsValue.getValue().isZero())
1733OpFoldResult DivUOp::fold(FoldAdaptor adaptor) {
1736 return foldDiv<
DivUOp,
false>(*
this, adaptor.getOperands());
1739OpFoldResult DivSOp::fold(FoldAdaptor adaptor) {
1745template <
class Op,
bool isSigned>
1746static OpFoldResult
foldMod(Op op, ArrayRef<Attribute> constants) {
1747 if (
auto rhsValue = dyn_cast_or_null<IntegerAttr>(constants[1])) {
1749 if (rhsValue.getValue() == 1)
1750 return getIntAttr(APInt::getZero(op.getType().getIntOrFloatBitWidth()),
1754 if (rhsValue.getValue().isZero())
1758 if (
auto lhsValue = dyn_cast_or_null<IntegerAttr>(constants[0])) {
1760 if (lhsValue.getValue().isZero())
1761 return getIntAttr(APInt::getZero(op.getType().getIntOrFloatBitWidth()),
1768OpFoldResult ModUOp::fold(FoldAdaptor adaptor) {
1771 return foldMod<
ModUOp,
false>(*
this, adaptor.getOperands());
1774OpFoldResult ModSOp::fold(FoldAdaptor adaptor) {
1780LogicalResult DivUOp::canonicalize(
DivUOp op, PatternRewriter &rewriter) {
1786LogicalResult ModUOp::canonicalize(
ModUOp op, PatternRewriter &rewriter) {
1798OpFoldResult ConcatOp::fold(FoldAdaptor adaptor) {
1802 if (getNumOperands() == 1)
1803 return getOperand(0);
1806 for (
auto attr : adaptor.getInputs())
1807 if (!attr || !isa<IntegerAttr>(attr))
1811 unsigned resultWidth = getType().getIntOrFloatBitWidth();
1812 APInt result(resultWidth, 0);
1814 unsigned nextInsertion = resultWidth;
1816 for (
auto attr : adaptor.getInputs()) {
1817 auto chunk = cast<IntegerAttr>(attr).getValue();
1818 nextInsertion -= chunk.getBitWidth();
1819 result.insertBits(chunk, nextInsertion);
1825LogicalResult ConcatOp::canonicalize(
ConcatOp op, PatternRewriter &rewriter) {
1829 auto inputs = op.getInputs();
1830 auto size = inputs.size();
1831 assert(size > 1 &&
"expected 2 or more operands");
1834 SmallVector<Value, 4> pendingOperands, processedOperands;
1835 bool anyOperandChanged;
1837 auto pushPendingOperands = [&](ValueRange operands) {
1838 auto size = operands.size();
1839 for (
size_t i = 0; i != size; ++i)
1840 pendingOperands.push_back(operands[size - 1 - i]);
1841 anyOperandChanged =
true;
1843 auto replacePrevOperand = [&](Value replacement) {
1844 processedOperands.back() = replacement;
1845 anyOperandChanged =
true;
1847 pendingOperands.reserve(size);
1848 pushPendingOperands(inputs);
1849 anyOperandChanged =
false;
1851 while (!pendingOperands.empty()) {
1852 Value nextOperand = pendingOperands.pop_back_val();
1856 if (
auto subConcat = nextOperand.getDefiningOp<
ConcatOp>()) {
1857 pushPendingOperands(subConcat->getOperands());
1862 if (!processedOperands.empty()) {
1863 Value prevOperand = processedOperands.back();
1867 if (
auto prevCst = prevOperand.getDefiningOp<
hw::ConstantOp>()) {
1868 unsigned prevWidth = prevCst.getValue().getBitWidth();
1869 unsigned thisWidth = cst.getValue().getBitWidth();
1870 auto resultCst = cst.getValue().zext(prevWidth + thisWidth);
1871 resultCst |= prevCst.getValue().zext(prevWidth + thisWidth)
1875 replacePrevOperand(replacement);
1881 if (nextOperand == prevOperand) {
1883 rewriter.createOrFold<ReplicateOp>(op.getLoc(), prevOperand, 2);
1884 replacePrevOperand(replacement);
1890 if (
auto repl = nextOperand.getDefiningOp<ReplicateOp>()) {
1892 if (repl.getOperand() == prevOperand) {
1893 Value replacement = rewriter.createOrFold<ReplicateOp>(
1894 op.getLoc(), repl.getOperand(), repl.getMultiple() + 1);
1895 replacePrevOperand(replacement);
1899 if (
auto prevRepl = prevOperand.getDefiningOp<ReplicateOp>()) {
1900 if (prevRepl.getOperand() == repl.getOperand()) {
1901 Value replacement = rewriter.createOrFold<ReplicateOp>(
1902 op.getLoc(), repl.getOperand(),
1903 repl.getMultiple() + prevRepl.getMultiple());
1904 replacePrevOperand(replacement);
1911 if (
auto repl = prevOperand.getDefiningOp<ReplicateOp>()) {
1912 if (repl.getOperand() == nextOperand) {
1913 Value replacement = rewriter.createOrFold<ReplicateOp>(
1914 op.getLoc(), nextOperand, repl.getMultiple() + 1);
1915 replacePrevOperand(replacement);
1922 if (
auto extract = nextOperand.getDefiningOp<
ExtractOp>()) {
1923 if (
auto prevExtract = prevOperand.getDefiningOp<
ExtractOp>()) {
1924 if (extract.getInput() == prevExtract.getInput()) {
1925 auto thisWidth = cast<IntegerType>(extract.getType()).getWidth();
1926 if (prevExtract.getLowBit() == extract.getLowBit() + thisWidth) {
1927 auto prevWidth = prevExtract.getType().getIntOrFloatBitWidth();
1928 auto resType = rewriter.getIntegerType(thisWidth + prevWidth);
1931 extract.getInput(), extract.getLowBit());
1932 replacePrevOperand(replacement);
1946 static std::optional<ArraySlice>
get(Value value) {
1947 assert(isa<IntegerType>(value.getType()) &&
"expected integer type");
1949 return ArraySlice{arrayGet.getInput(), arrayGet.getIndex(), 1};
1952 if (
auto arraySlice =
1955 arraySlice.getInput(), arraySlice.getLowIndex(),
1956 hw::type_cast<hw::ArrayType>(arraySlice.getType())
1958 return std::nullopt;
1961 if (
auto extractOpt = ArraySlice::get(nextOperand)) {
1962 if (
auto prevExtractOpt = ArraySlice::get(prevOperand)) {
1964 if (prevExtractOpt->index.getType() == extractOpt->index.getType() &&
1965 prevExtractOpt->input == extractOpt->input &&
1966 hw::isOffset(extractOpt->index, prevExtractOpt->index,
1967 extractOpt->width)) {
1968 auto resType = hw::ArrayType::get(
1969 hw::type_cast<hw::ArrayType>(prevExtractOpt->input.getType())
1971 extractOpt->width + prevExtractOpt->width);
1972 auto resIntType = rewriter.getIntegerType(hw::getBitWidth(resType));
1974 rewriter, op.getLoc(), resIntType,
1976 prevExtractOpt->input,
1977 extractOpt->index));
1978 replacePrevOperand(replacement);
1985 processedOperands.push_back(nextOperand);
1997 constexpr size_t kBatchExtractThreshold = 16;
1998 bool anyExtractsResolved =
false;
1999 if (!anyOperandChanged && processedOperands.size() > 1) {
2000 SmallVector<ExtractOp, 8> extractUsers;
2001 for (
auto *user : op->getUsers()) {
2002 if (
auto extract = dyn_cast<ExtractOp>(user))
2003 extractUsers.push_back(extract);
2006 if (extractUsers.size() >= kBatchExtractThreshold) {
2008 auto concatInputs = op.getInputs();
2009 size_t numConcatOperands = concatInputs.size();
2010 SmallVector<size_t> prefixWidths(numConcatOperands);
2011 size_t cumWidth = 0;
2012 for (
size_t i = 0; i < numConcatOperands; ++i) {
2015 cumWidth += concatInputs[numConcatOperands - 1 - i]
2017 .getIntOrFloatBitWidth();
2018 prefixWidths[i] = cumWidth;
2022 for (
auto extract : extractUsers) {
2025 anyExtractsResolved =
true;
2030 if (processedOperands.size() == 1) {
2034 }
else if (anyOperandChanged) {
2035 replaceOpWithNewOpAndCopyNamehint<ConcatOp>(rewriter, op, op.getType(),
2037 }
else if (!anyExtractsResolved) {
2047OpFoldResult MuxOp::fold(FoldAdaptor adaptor) {
2052 if (getTrueValue() == getFalseValue() && getTrueValue() != getResult())
2053 return getTrueValue();
2054 if (
auto tv = adaptor.getTrueValue())
2055 if (tv == adaptor.getFalseValue())
2060 if (
auto pred = dyn_cast_or_null<IntegerAttr>(adaptor.getCond())) {
2061 if (pred.getValue().isZero() && getFalseValue() != getResult())
2062 return getFalseValue();
2063 if (pred.getValue().isOne() && getTrueValue() != getResult())
2064 return getTrueValue();
2068 if (getCond().getType() == getTrueValue().getType())
2069 if (
auto tv = dyn_cast_or_null<IntegerAttr>(adaptor.getTrueValue()))
2070 if (
auto fv = dyn_cast_or_null<IntegerAttr>(adaptor.getFalseValue()))
2071 if (tv.getValue().isOne() && fv.getValue().isZero() &&
2072 hw::getBitWidth(getType()) == 1 && getCond() != getResult())
2088 if (
auto cmp = cond.getDefiningOp<ICmpOp>()) {
2090 auto requiredPredicate =
2091 (isInverted ? ICmpPredicate::eq : ICmpPredicate::ne);
2092 if (cmp.getLhs() == indexValue && cmp.getPredicate() == requiredPredicate) {
2102 if (
auto orOp = cond.getDefiningOp<
OrOp>()) {
2105 for (
auto operand : orOp.getOperands())
2112 if (
auto andOp = cond.getDefiningOp<
AndOp>()) {
2115 for (
auto operand : andOp.getOperands())
2134 PatternRewriter &rewriter,
MuxOp rootMux,
bool isFalseSide,
2140 auto rootCmp = rootMux.getCond().getDefiningOp<ICmpOp>();
2143 Value indexValue = rootCmp.getLhs();
2146 auto getCaseValue = [&](
MuxOp mux) -> Value {
2147 return mux.getOperand(1 +
unsigned(!isFalseSide));
2152 auto getTreeValue = [&](
MuxOp mux) -> Value {
2153 return mux.getOperand(1 +
unsigned(isFalseSide));
2158 SmallVector<Location> locationsFound;
2159 SmallVector<std::pair<hw::ConstantOp, Value>, 4> valuesFound;
2163 auto collectConstantValues = [&](
MuxOp mux) ->
bool {
2165 mux.getCond(), indexValue, isFalseSide, [&](
hw::ConstantOp cst) {
2166 valuesFound.push_back({cst, getCaseValue(mux)});
2167 locationsFound.push_back(mux.getCond().getLoc());
2168 locationsFound.push_back(mux->getLoc());
2173 if (!collectConstantValues(rootMux))
2177 if (rootMux->hasOneUse()) {
2178 if (
auto userMux = dyn_cast<MuxOp>(*rootMux->user_begin())) {
2179 if (getTreeValue(userMux) == rootMux.getResult() &&
2187 auto nextTreeValue = getTreeValue(rootMux);
2189 auto nextMux = nextTreeValue.getDefiningOp<
MuxOp>();
2190 if (!nextMux || !nextMux->hasOneUse())
2192 if (!collectConstantValues(nextMux))
2194 nextTreeValue = getTreeValue(nextMux);
2197 auto indexWidth = cast<IntegerType>(indexValue.getType()).getWidth();
2199 if (indexWidth > 20)
2202 auto foldingStyle = styleFn(indexWidth, valuesFound.size());
2206 uint64_t tableSize = 1ULL << indexWidth;
2210 SmallVector<Value, 8> table(tableSize, nextTreeValue);
2215 for (
auto &elt :
llvm::reverse(valuesFound)) {
2216 uint64_t idx = elt.first.getValue().getZExtValue();
2217 assert(idx < table.size() &&
"constant should be same bitwidth as index");
2218 table[idx] = elt.second;
2222 SmallVector<Value> bits;
2231 "unknown folding style");
2235 std::reverse(table.begin(), table.end());
2238 auto fusedLoc = rewriter.getFusedLoc(locationsFound);
2240 replaceOpWithNewOpAndCopyNamehint<hw::ArrayGetOp>(rewriter, rootMux, array,
2255 PatternRewriter &rewriter) {
2256 assert(fullyAssoc->getNumOperands() >= 2 &&
"cannot split up unary ops");
2257 assert(operandNo < fullyAssoc->getNumOperands() &&
"Invalid operand #");
2261 if (fullyAssoc->getNumOperands() == 2)
2262 return fullyAssoc->getOperand(operandNo ^ 1);
2265 if (fullyAssoc->hasOneUse()) {
2266 rewriter.modifyOpInPlace(fullyAssoc,
2267 [&]() { fullyAssoc->eraseOperand(operandNo); });
2268 return fullyAssoc->getResult(0);
2272 SmallVector<Value> operands;
2273 operands.append(fullyAssoc->getOperands().begin(),
2274 fullyAssoc->getOperands().begin() + operandNo);
2275 operands.append(fullyAssoc->getOperands().begin() + operandNo + 1,
2276 fullyAssoc->getOperands().end());
2278 fullyAssoc->getLoc(), fullyAssoc->getName(), operands, rewriter);
2279 Value excluded = fullyAssoc->getOperand(operandNo);
2283 ArrayRef<Value>{opWithoutExcluded, excluded}, rewriter);
2285 return opWithoutExcluded;
2295 PatternRewriter &rewriter) {
2298 Operation *subExpr =
2299 (isTrueOperand ? op.getFalseValue() : op.getTrueValue()).getDefiningOp();
2300 if (!subExpr || subExpr->getNumOperands() < 2)
2304 if (!isa<AndOp, XorOp, OrOp, MuxOp>(subExpr))
2309 Value commonValue = isTrueOperand ? op.getTrueValue() : op.getFalseValue();
2310 size_t opNo = 0, e = subExpr->getNumOperands();
2311 while (opNo != e && subExpr->getOperand(opNo) != commonValue)
2317 Value cond = op.getCond();
2323 if (
auto subMux = dyn_cast<MuxOp>(subExpr)) {
2328 Value subCond = subMux.getCond();
2331 if (subMux.getTrueValue() == commonValue)
2332 otherValue = subMux.getFalseValue();
2333 else if (subMux.getFalseValue() == commonValue) {
2334 otherValue = subMux.getTrueValue();
2344 cond = rewriter.createOrFold<
OrOp>(op.getLoc(), cond, subCond,
false);
2345 replaceOpWithNewOpAndCopyNamehint<MuxOp>(rewriter, op, cond, commonValue,
2346 otherValue, op.getTwoState());
2352 bool isaAndOp = isa<AndOp>(subExpr);
2353 if (isTrueOperand ^ isaAndOp)
2357 rewriter.createOrFold<ReplicateOp>(op.getLoc(), op.getType(), cond);
2360 bool isaXorOp = isa<XorOp>(subExpr);
2361 bool isaOrOp = isa<OrOp>(subExpr);
2370 if (isaOrOp || isaXorOp) {
2371 auto masked = rewriter.createOrFold<
AndOp>(op.getLoc(), extendedCond,
2372 restOfAssoc,
false);
2374 replaceOpWithNewOpAndCopyNamehint<XorOp>(rewriter, op, masked,
2375 commonValue,
false);
2377 replaceOpWithNewOpAndCopyNamehint<OrOp>(rewriter, op, masked, commonValue,
2383 assert(isaAndOp &&
"unexpected operation here");
2384 auto masked = rewriter.createOrFold<
OrOp>(op.getLoc(), extendedCond,
2385 restOfAssoc,
false);
2386 replaceOpWithNewOpAndCopyNamehint<AndOp>(rewriter, op, masked, commonValue,
2397 PatternRewriter &rewriter) {
2400 if (!isa<ConcatOp>(trueOp))
2404 SmallVector<Value> trueOperands, falseOperands;
2408 size_t numTrueOperands = trueOperands.size();
2409 size_t numFalseOperands = falseOperands.size();
2411 if (!numTrueOperands || !numFalseOperands ||
2412 (trueOperands.front() != falseOperands.front() &&
2413 trueOperands.back() != falseOperands.back()))
2417 if (trueOperands.front() == falseOperands.front()) {
2418 SmallVector<Value> operands;
2420 for (i = 0; i < numTrueOperands; ++i) {
2421 Value trueOperand = trueOperands[i];
2422 if (trueOperand == falseOperands[i])
2423 operands.push_back(trueOperand);
2427 if (i == numTrueOperands) {
2434 if (llvm::all_of(operands, [&](Value v) {
return v == operands.front(); }))
2435 sharedMSB = rewriter.createOrFold<ReplicateOp>(
2436 mux->getLoc(), operands.front(), operands.size());
2438 sharedMSB = rewriter.createOrFold<
ConcatOp>(mux->getLoc(), operands);
2442 operands.append(trueOperands.begin() + i, trueOperands.end());
2443 Value trueLSB = rewriter.createOrFold<
ConcatOp>(trueOp->getLoc(), operands);
2445 operands.append(falseOperands.begin() + i, falseOperands.end());
2447 rewriter.createOrFold<
ConcatOp>(falseOp->getLoc(), operands);
2450 Value lsb = rewriter.createOrFold<
MuxOp>(
2451 mux->getLoc(), mux.getCond(), trueLSB, falseLSB, mux.getTwoState());
2452 replaceOpWithNewOpAndCopyNamehint<ConcatOp>(rewriter, mux, sharedMSB, lsb);
2457 if (trueOperands.back() == falseOperands.back()) {
2458 SmallVector<Value> operands;
2461 Value trueOperand = trueOperands[numTrueOperands - i - 1];
2462 if (trueOperand == falseOperands[numFalseOperands - i - 1])
2463 operands.push_back(trueOperand);
2467 std::reverse(operands.begin(), operands.end());
2468 Value sharedLSB = rewriter.createOrFold<
ConcatOp>(mux->getLoc(), operands);
2472 operands.append(trueOperands.begin(), trueOperands.end() - i);
2473 Value trueMSB = rewriter.createOrFold<
ConcatOp>(trueOp->getLoc(), operands);
2475 operands.append(falseOperands.begin(), falseOperands.end() - i);
2477 rewriter.createOrFold<
ConcatOp>(falseOp->getLoc(), operands);
2479 Value msb = rewriter.createOrFold<
MuxOp>(
2480 mux->getLoc(), mux.getCond(), trueMSB, falseMSB, mux.getTwoState());
2481 replaceOpWithNewOpAndCopyNamehint<ConcatOp>(rewriter, mux, msb, sharedLSB);
2493 if (!trueVec || !falseVec)
2495 if (!trueVec.isUniform() || !falseVec.isUniform())
2498 auto mux = MuxOp::create(rewriter, op.getLoc(), op.getCond(),
2499 trueVec.getUniformElement(),
2500 falseVec.getUniformElement(), op.getTwoState());
2502 SmallVector<Value> values(trueVec.getInputs().size(), mux);
2510 bool constCond, PatternRewriter &rewriter) {
2511 if (!muxValue.hasOneUse())
2513 auto *op = muxValue.getDefiningOp();
2514 if (!op || !isa_and_nonnull<CombDialect>(op->getDialect()))
2516 if (!llvm::is_contained(op->getOperands(), muxCond))
2518 OpBuilder::InsertionGuard guard(rewriter);
2519 rewriter.setInsertionPoint(op);
2522 rewriter.modifyOpInPlace(op, [&] {
2523 for (
auto &use : op->getOpOperands())
2524 if (use.get() == muxCond)
2532 using OpRewritePattern::OpRewritePattern;
2534 LogicalResult matchAndRewrite(
MuxOp op,
2535 PatternRewriter &rewriter)
const override;
2539foldToArrayCreateOnlyWhenDense(
size_t indexWidth,
size_t numEntries) {
2542 if (indexWidth >= 9 || numEntries < 3)
2548 uint64_t tableSize = 1ULL << indexWidth;
2549 if (numEntries >= tableSize * 5 / 8)
2554LogicalResult MuxRewriter::matchAndRewrite(
MuxOp op,
2555 PatternRewriter &rewriter)
const {
2559 bool isSignlessInt =
false;
2560 if (
auto intType = dyn_cast<IntegerType>(op.getType()))
2561 isSignlessInt = intType.isSignless();
2568 if (matchPattern(op.getTrueValue(), m_ConstantInt(&value)) && isSignlessInt) {
2569 if (value.getBitWidth() == 1) {
2571 if (value.isZero()) {
2573 replaceOpWithNewOpAndCopyNamehint<AndOp>(rewriter, op, notCond,
2574 op.getFalseValue(),
false);
2579 replaceOpWithNewOpAndCopyNamehint<OrOp>(rewriter, op, op.getCond(),
2580 op.getFalseValue(),
false);
2586 if (matchPattern(op.getFalseValue(), m_ConstantInt(&value2))) {
2591 APInt xorValue = value ^ value2;
2592 if (xorValue.isPowerOf2()) {
2593 unsigned leadingZeros = xorValue.countLeadingZeros();
2594 unsigned trailingZeros = value.getBitWidth() - leadingZeros - 1;
2595 SmallVector<Value, 3> operands;
2603 if (leadingZeros > 0)
2604 operands.push_back(rewriter.createOrFold<
ExtractOp>(
2605 op.getLoc(), op.getTrueValue(), trailingZeros + 1, leadingZeros));
2609 auto v1 = rewriter.createOrFold<
ExtractOp>(
2610 op.getLoc(), op.getTrueValue(), trailingZeros, 1);
2611 auto v2 = rewriter.createOrFold<
ExtractOp>(
2612 op.getLoc(), op.getFalseValue(), trailingZeros, 1);
2613 operands.push_back(rewriter.createOrFold<
MuxOp>(
2614 op.getLoc(), op.getCond(), v1, v2,
false));
2616 if (trailingZeros > 0)
2617 operands.push_back(rewriter.createOrFold<
ExtractOp>(
2618 op.getLoc(), op.getTrueValue(), 0, trailingZeros));
2620 replaceOpWithNewOpAndCopyNamehint<ConcatOp>(rewriter, op, op.getType(),
2627 if (value.isAllOnes() && value2.isZero()) {
2628 replaceOpWithNewOpAndCopyNamehint<ReplicateOp>(
2629 rewriter, op, op.getType(), op.getCond());
2635 if (matchPattern(op.getFalseValue(), m_ConstantInt(&value)) &&
2636 isSignlessInt && value.getBitWidth() == 1) {
2638 if (value.isZero()) {
2639 replaceOpWithNewOpAndCopyNamehint<AndOp>(rewriter, op, op.getCond(),
2640 op.getTrueValue(),
false);
2647 auto notCond = rewriter.createOrFold<
XorOp>(op.getLoc(), op.getCond(),
2648 op.getFalseValue(),
false);
2649 replaceOpWithNewOpAndCopyNamehint<OrOp>(rewriter, op, notCond,
2650 op.getTrueValue(),
false);
2656 Operation *condOp = op.getCond().getDefiningOp();
2657 if (condOp && matchPattern(condOp,
m_Complement(m_Any(&subExpr))) &&
2659 replaceOpWithNewOpAndCopyNamehint<MuxOp>(rewriter, op, op.getType(),
2660 subExpr, op.getFalseValue(),
2661 op.getTrueValue(),
true);
2668 if (condOp && condOp->hasOneUse()) {
2669 SmallVector<Value> invertedOperands;
2673 auto getInvertedOperands = [&]() ->
bool {
2674 for (Value operand : condOp->getOperands()) {
2675 if (matchPattern(operand,
m_Complement(m_Any(&subExpr))))
2676 invertedOperands.push_back(subExpr);
2683 if (isa<AndOp>(condOp) && getInvertedOperands()) {
2685 rewriter.createOrFold<
OrOp>(op.getLoc(), invertedOperands,
false);
2686 replaceOpWithNewOpAndCopyNamehint<MuxOp>(
2687 rewriter, op, newOr, op.getFalseValue(), op.getTrueValue(),
2691 if (isa<OrOp>(condOp) && getInvertedOperands()) {
2693 rewriter.createOrFold<
AndOp>(op.getLoc(), invertedOperands,
false);
2694 replaceOpWithNewOpAndCopyNamehint<MuxOp>(
2695 rewriter, op, newAnd, op.getFalseValue(), op.getTrueValue(),
2701 if (
auto falseMux = op.getFalseValue().getDefiningOp<
MuxOp>();
2702 falseMux && falseMux != op) {
2704 if (op.getCond() == falseMux.getCond() &&
2705 falseMux.getFalseValue() != falseMux) {
2706 replaceOpWithNewOpAndCopyNamehint<MuxOp>(
2707 rewriter, op, op.getCond(), op.getTrueValue(),
2708 falseMux.getFalseValue(), op.getTwoStateAttr());
2714 foldToArrayCreateOnlyWhenDense))
2718 if (
auto trueMux = op.getTrueValue().getDefiningOp<
MuxOp>();
2719 trueMux && trueMux != op) {
2721 if (op.getCond() == trueMux.getCond()) {
2722 replaceOpWithNewOpAndCopyNamehint<MuxOp>(
2723 rewriter, op, op.getCond(), trueMux.getTrueValue(),
2724 op.getFalseValue(), op.getTwoStateAttr());
2730 foldToArrayCreateOnlyWhenDense))
2735 if (
auto trueMux = dyn_cast_or_null<MuxOp>(op.getTrueValue().getDefiningOp()),
2736 falseMux = dyn_cast_or_null<MuxOp>(op.getFalseValue().getDefiningOp());
2737 trueMux && falseMux && trueMux.getCond() == falseMux.getCond() &&
2738 trueMux.getTrueValue() == falseMux.getTrueValue() && trueMux != op &&
2740 auto subMux = MuxOp::create(
2741 rewriter, rewriter.getFusedLoc({trueMux.getLoc(), falseMux.getLoc()}),
2742 op.getCond(), trueMux.getFalseValue(), falseMux.getFalseValue());
2743 replaceOpWithNewOpAndCopyNamehint<MuxOp>(rewriter, op, trueMux.getCond(),
2744 trueMux.getTrueValue(), subMux,
2745 op.getTwoStateAttr());
2750 if (
auto trueMux = dyn_cast_or_null<MuxOp>(op.getTrueValue().getDefiningOp()),
2751 falseMux = dyn_cast_or_null<MuxOp>(op.getFalseValue().getDefiningOp());
2752 trueMux && falseMux && trueMux.getCond() == falseMux.getCond() &&
2753 trueMux.getFalseValue() == falseMux.getFalseValue() && trueMux != op &&
2755 auto subMux = MuxOp::create(
2756 rewriter, rewriter.getFusedLoc({trueMux.getLoc(), falseMux.getLoc()}),
2757 op.getCond(), trueMux.getTrueValue(), falseMux.getTrueValue());
2758 replaceOpWithNewOpAndCopyNamehint<MuxOp>(rewriter, op, trueMux.getCond(),
2759 subMux, trueMux.getFalseValue(),
2760 op.getTwoStateAttr());
2765 if (
auto trueMux = dyn_cast_or_null<MuxOp>(op.getTrueValue().getDefiningOp()),
2766 falseMux = dyn_cast_or_null<MuxOp>(op.getFalseValue().getDefiningOp());
2767 trueMux && falseMux &&
2768 trueMux.getTrueValue() == falseMux.getTrueValue() &&
2769 trueMux.getFalseValue() == falseMux.getFalseValue() && trueMux != op &&
2772 MuxOp::create(rewriter,
2773 rewriter.getFusedLoc(
2774 {op.getLoc(), trueMux.getLoc(), falseMux.getLoc()}),
2775 op.getCond(), trueMux.getCond(), falseMux.getCond());
2776 replaceOpWithNewOpAndCopyNamehint<MuxOp>(
2777 rewriter, op, subMux, trueMux.getTrueValue(), trueMux.getFalseValue(),
2778 op.getTwoStateAttr());
2790 if (Operation *trueOp = op.getTrueValue().getDefiningOp())
2791 if (Operation *falseOp = op.getFalseValue().getDefiningOp())
2792 if (trueOp->getName() == falseOp->getName())
2805 if (op.getTrueValue().getDefiningOp() &&
2806 op.getTrueValue().getDefiningOp() != op)
2809 if (op.getFalseValue().getDefiningOp() &&
2810 op.getFalseValue().getDefiningOp() != op)
2821 if (op.getInputs().empty() || op.isUniform())
2823 auto inputs = op.getInputs();
2824 if (inputs.size() <= 1)
2829 auto first = inputs[0].getDefiningOp<
comb::MuxOp>();
2834 for (
size_t i = 1, n = inputs.size(); i < n; ++i) {
2835 auto input = inputs[i].getDefiningOp<
comb::MuxOp>();
2836 if (!input || first.getCond() != input.getCond())
2841 SmallVector<Value> trues{first.getTrueValue()};
2842 SmallVector<Value> falses{first.getFalseValue()};
2843 SmallVector<Location> locs{first->getLoc()};
2844 bool isTwoState =
true;
2845 for (
size_t i = 1, n = inputs.size(); i < n; ++i) {
2846 auto input = inputs[i].getDefiningOp<
comb::MuxOp>();
2847 trues.push_back(input.getTrueValue());
2848 falses.push_back(input.getFalseValue());
2849 locs.push_back(input->getLoc());
2850 if (!input.getTwoState())
2855 auto loc = FusedLoc::get(op.getContext(), locs);
2859 auto arrayTy = op.getType();
2862 rewriter.replaceOpWithNewOp<
comb::MuxOp>(op, arrayTy, first.getCond(),
2863 trueValues, falseValues, isTwoState);
2868 using OpRewritePattern::OpRewritePattern;
2871 PatternRewriter &rewriter)
const override {
2872 if (foldArrayOfMuxes(op, rewriter))
2880void MuxOp::getCanonicalizationPatterns(RewritePatternSet &results,
2882 results.insert<MuxRewriter, ArrayRewriter>(
context);
2893 switch (predicate) {
2894 case ICmpPredicate::eq:
2896 case ICmpPredicate::ne:
2898 case ICmpPredicate::slt:
2899 return lhs.slt(rhs);
2900 case ICmpPredicate::sle:
2901 return lhs.sle(rhs);
2902 case ICmpPredicate::sgt:
2903 return lhs.sgt(rhs);
2904 case ICmpPredicate::sge:
2905 return lhs.sge(rhs);
2906 case ICmpPredicate::ult:
2907 return lhs.ult(rhs);
2908 case ICmpPredicate::ule:
2909 return lhs.ule(rhs);
2910 case ICmpPredicate::ugt:
2911 return lhs.ugt(rhs);
2912 case ICmpPredicate::uge:
2913 return lhs.uge(rhs);
2914 case ICmpPredicate::ceq:
2916 case ICmpPredicate::cne:
2918 case ICmpPredicate::weq:
2920 case ICmpPredicate::wne:
2923 llvm_unreachable(
"unknown comparison predicate");
2929 switch (predicate) {
2930 case ICmpPredicate::eq:
2931 case ICmpPredicate::sle:
2932 case ICmpPredicate::sge:
2933 case ICmpPredicate::ule:
2934 case ICmpPredicate::uge:
2935 case ICmpPredicate::ceq:
2936 case ICmpPredicate::weq:
2938 case ICmpPredicate::ne:
2939 case ICmpPredicate::slt:
2940 case ICmpPredicate::sgt:
2941 case ICmpPredicate::ult:
2942 case ICmpPredicate::ugt:
2943 case ICmpPredicate::cne:
2944 case ICmpPredicate::wne:
2947 llvm_unreachable(
"unknown comparison predicate");
2950OpFoldResult ICmpOp::fold(FoldAdaptor adaptor) {
2953 if (getLhs() == getRhs()) {
2955 return IntegerAttr::get(getType(), val);
2959 if (
auto lhs = dyn_cast_or_null<IntegerAttr>(adaptor.getLhs())) {
2960 if (
auto rhs = dyn_cast_or_null<IntegerAttr>(adaptor.getRhs())) {
2963 return IntegerAttr::get(getType(), val);
2971template <
typename Range>
2973 size_t commonPrefixLength = 0;
2974 auto ia = a.begin();
2975 auto ib = b.begin();
2977 for (; ia != a.end() && ib != b.end(); ia++, ib++, commonPrefixLength++) {
2983 return commonPrefixLength;
2987 size_t totalWidth = 0;
2988 for (
auto operand : operands) {
2991 ssize_t width = operand.getType().getIntOrFloatBitWidth();
2993 totalWidth += width;
3003 PatternRewriter &rewriter) {
3007 SmallVector<Value> lhsOperands, rhsOperands;
3010 ArrayRef<Value> lhsOperandsRef = lhsOperands, rhsOperandsRef = rhsOperands;
3012 auto formCatOrReplicate = [&](Location loc,
3013 ArrayRef<Value> operands) -> Value {
3014 assert(!operands.empty());
3015 Value sameElement = operands[0];
3016 for (
size_t i = 1, e = operands.size(); i != e && sameElement; ++i)
3017 if (sameElement != operands[i])
3018 sameElement = Value();
3020 return rewriter.createOrFold<ReplicateOp>(loc, sameElement,
3022 return rewriter.createOrFold<
ConcatOp>(loc, operands);
3025 auto replaceWith = [&](ICmpPredicate predicate, Value lhs,
3026 Value rhs) -> LogicalResult {
3027 replaceOpWithNewOpAndCopyNamehint<ICmpOp>(rewriter, op, predicate, lhs, rhs,
3032 size_t commonPrefixLength =
3034 if (commonPrefixLength == lhsOperands.size()) {
3037 replaceOpWithNewOpAndCopyNamehint<hw::ConstantOp>(rewriter, op,
3043 llvm::reverse(lhsOperandsRef), llvm::reverse(rhsOperandsRef));
3045 size_t commonPrefixTotalWidth =
3046 getTotalWidth(lhsOperandsRef.take_front(commonPrefixLength));
3047 size_t commonSuffixTotalWidth =
3048 getTotalWidth(lhsOperandsRef.take_back(commonSuffixLength));
3049 auto lhsOnly = lhsOperandsRef.drop_front(commonPrefixLength)
3050 .drop_back(commonSuffixLength);
3051 auto rhsOnly = rhsOperandsRef.drop_front(commonPrefixLength)
3052 .drop_back(commonSuffixLength);
3054 auto replaceWithoutReplicatingSignBit = [&]() {
3055 auto newLhs = formCatOrReplicate(lhs->getLoc(), lhsOnly);
3056 auto newRhs = formCatOrReplicate(rhs->getLoc(), rhsOnly);
3057 return replaceWith(op.getPredicate(), newLhs, newRhs);
3060 auto replaceWithReplicatingSignBit = [&]() {
3061 auto firstNonEmptyValue = lhsOperands[0];
3062 auto firstNonEmptyElemWidth =
3063 firstNonEmptyValue.getType().getIntOrFloatBitWidth();
3064 Value signBit = rewriter.createOrFold<
ExtractOp>(
3065 op.getLoc(), firstNonEmptyValue, firstNonEmptyElemWidth - 1, 1);
3067 auto newLhs = ConcatOp::create(rewriter, lhs->getLoc(), signBit, lhsOnly);
3068 auto newRhs = ConcatOp::create(rewriter, rhs->getLoc(), signBit, rhsOnly);
3069 return replaceWith(op.getPredicate(), newLhs, newRhs);
3072 if (ICmpOp::isPredicateSigned(op.getPredicate())) {
3074 if (commonPrefixTotalWidth == 0 && commonSuffixTotalWidth > 0)
3075 return replaceWithoutReplicatingSignBit();
3081 if (commonPrefixTotalWidth > 1 || commonSuffixTotalWidth > 0)
3082 return replaceWithReplicatingSignBit();
3084 }
else if (commonPrefixTotalWidth > 0 || commonSuffixTotalWidth > 0) {
3086 return replaceWithoutReplicatingSignBit();
3100 ICmpOp cmpOp,
const KnownBits &bitAnalysis,
const APInt &rhsCst,
3101 PatternRewriter &rewriter) {
3105 APInt bitsKnown = bitAnalysis.Zero | bitAnalysis.One;
3106 if ((bitsKnown & rhsCst) != bitAnalysis.One) {
3109 bool result = cmpOp.getPredicate() == ICmpPredicate::ne;
3110 replaceOpWithNewOpAndCopyNamehint<hw::ConstantOp>(rewriter, cmpOp,
3118 SmallVector<Value> newConcatOperands;
3119 auto newConstant = APInt::getZeroWidth();
3124 unsigned knownMSB = bitsKnown.countLeadingOnes();
3126 Value operand = cmpOp.getLhs();
3131 while (knownMSB != bitsKnown.getBitWidth()) {
3134 bitsKnown = bitsKnown.trunc(bitsKnown.getBitWidth() - knownMSB);
3137 unsigned unknownBits = bitsKnown.countLeadingZeros();
3138 unsigned lowBit = bitsKnown.getBitWidth() - unknownBits;
3139 auto spanOperand = rewriter.createOrFold<
ExtractOp>(
3140 operand.getLoc(), operand, lowBit,
3142 auto spanConstant = rhsCst.lshr(lowBit).trunc(unknownBits);
3145 newConcatOperands.push_back(spanOperand);
3148 if (newConstant.getBitWidth() != 0)
3149 newConstant = newConstant.concat(spanConstant);
3151 newConstant = spanConstant;
3154 unsigned newWidth = bitsKnown.getBitWidth() - unknownBits;
3155 bitsKnown = bitsKnown.trunc(newWidth);
3156 knownMSB = bitsKnown.countLeadingOnes();
3162 if (newConcatOperands.empty()) {
3163 bool result = cmpOp.getPredicate() == ICmpPredicate::eq;
3164 replaceOpWithNewOpAndCopyNamehint<hw::ConstantOp>(rewriter, cmpOp,
3170 Value concatResult =
3171 rewriter.createOrFold<
ConcatOp>(operand.getLoc(), newConcatOperands);
3175 rewriter, cmpOp.getOperand(1).getLoc(), newConstant);
3177 replaceOpWithNewOpAndCopyNamehint<ICmpOp>(rewriter, cmpOp,
3178 cmpOp.getPredicate(), concatResult,
3179 newConstantOp, cmpOp.getTwoState());
3185 PatternRewriter &rewriter) {
3186 auto ip = rewriter.saveInsertionPoint();
3187 rewriter.setInsertionPoint(xorOp);
3189 auto xorRHS = xorOp.getOperands().back().getDefiningOp<
hw::ConstantOp>();
3191 xorRHS.getValue() ^ rhs);
3193 switch (xorOp.getNumOperands()) {
3197 APInt::getZero(rhs.getBitWidth()));
3201 newLHS = xorOp.getOperand(0);
3205 SmallVector<Value> newOperands(xorOp.getOperands());
3206 newOperands.pop_back();
3207 newLHS = XorOp::create(rewriter, xorOp.getLoc(), newOperands,
false);
3211 bool xorMultipleUses = !xorOp->hasOneUse();
3215 if (xorMultipleUses)
3216 replaceOpWithNewOpAndCopyNamehint<XorOp>(rewriter, xorOp, newLHS, xorRHS,
3220 rewriter.restoreInsertionPoint(ip);
3221 replaceOpWithNewOpAndCopyNamehint<ICmpOp>(
3222 rewriter, cmpOp, cmpOp.getPredicate(), newLHS, newRHS,
false);
3225LogicalResult ICmpOp::canonicalize(ICmpOp op, PatternRewriter &rewriter) {
3231 if (matchPattern(op.getLhs(), m_ConstantInt(&lhs))) {
3232 assert(!matchPattern(op.getRhs(), m_ConstantInt(&rhs)) &&
3233 "Should be folded");
3234 replaceOpWithNewOpAndCopyNamehint<ICmpOp>(
3235 rewriter, op, ICmpOp::getFlippedPredicate(op.getPredicate()),
3236 op.getRhs(), op.getLhs(), op.getTwoState());
3241 if (matchPattern(op.getRhs(), m_ConstantInt(&rhs))) {
3246 auto replaceWith = [&](ICmpPredicate predicate, Value lhs,
3247 Value rhs) -> LogicalResult {
3248 replaceOpWithNewOpAndCopyNamehint<ICmpOp>(rewriter, op, predicate, lhs,
3249 rhs, op.getTwoState());
3253 auto replaceWithConstantI1 = [&](
bool constant) -> LogicalResult {
3254 replaceOpWithNewOpAndCopyNamehint<hw::ConstantOp>(rewriter, op,
3255 APInt(1, constant));
3259 switch (op.getPredicate()) {
3260 case ICmpPredicate::slt:
3262 if (rhs.isMaxSignedValue())
3263 return replaceWith(ICmpPredicate::ne, op.getLhs(), op.getRhs());
3265 if (rhs.isMinSignedValue())
3266 return replaceWithConstantI1(0);
3268 if ((rhs - 1).isMinSignedValue())
3269 return replaceWith(ICmpPredicate::eq, op.getLhs(),
3272 case ICmpPredicate::sgt:
3274 if (rhs.isMinSignedValue())
3275 return replaceWith(ICmpPredicate::ne, op.getLhs(), op.getRhs());
3277 if (rhs.isMaxSignedValue())
3278 return replaceWithConstantI1(0);
3280 if ((rhs + 1).isMaxSignedValue())
3281 return replaceWith(ICmpPredicate::eq, op.getLhs(),
3284 case ICmpPredicate::ult:
3286 if (rhs.isAllOnes())
3287 return replaceWith(ICmpPredicate::ne, op.getLhs(), op.getRhs());
3290 return replaceWithConstantI1(0);
3292 if ((rhs - 1).isZero())
3293 return replaceWith(ICmpPredicate::eq, op.getLhs(),
3297 if (rhs.countLeadingOnes() + rhs.countTrailingZeros() ==
3298 rhs.getBitWidth()) {
3299 auto numOnes = rhs.countLeadingOnes();
3301 rhs.getBitWidth() - numOnes, numOnes);
3302 return replaceWith(ICmpPredicate::ne, smaller,
3307 case ICmpPredicate::ugt:
3310 return replaceWith(ICmpPredicate::ne, op.getLhs(), op.getRhs());
3312 if (rhs.isAllOnes())
3313 return replaceWithConstantI1(0);
3315 if ((rhs + 1).isAllOnes())
3316 return replaceWith(ICmpPredicate::eq, op.getLhs(),
3320 if ((rhs + 1).isPowerOf2()) {
3321 auto numOnes = rhs.countTrailingOnes();
3322 auto newWidth = rhs.getBitWidth() - numOnes;
3325 return replaceWith(ICmpPredicate::ne, smaller,
3330 case ICmpPredicate::sle:
3332 if (rhs.isMaxSignedValue())
3333 return replaceWithConstantI1(1);
3335 return replaceWith(ICmpPredicate::slt, op.getLhs(),
getConstant(rhs + 1));
3336 case ICmpPredicate::sge:
3338 if (rhs.isMinSignedValue())
3339 return replaceWithConstantI1(1);
3341 return replaceWith(ICmpPredicate::sgt, op.getLhs(),
getConstant(rhs - 1));
3342 case ICmpPredicate::ule:
3344 if (rhs.isAllOnes())
3345 return replaceWithConstantI1(1);
3347 return replaceWith(ICmpPredicate::ult, op.getLhs(),
getConstant(rhs + 1));
3348 case ICmpPredicate::uge:
3351 return replaceWithConstantI1(1);
3353 return replaceWith(ICmpPredicate::ugt, op.getLhs(),
getConstant(rhs - 1));
3354 case ICmpPredicate::eq:
3355 if (rhs.getBitWidth() == 1) {
3358 replaceOpWithNewOpAndCopyNamehint<XorOp>(rewriter, op, op.getLhs(),
3363 if (rhs.isAllOnes()) {
3370 case ICmpPredicate::ne:
3371 if (rhs.getBitWidth() == 1) {
3377 if (rhs.isAllOnes()) {
3379 replaceOpWithNewOpAndCopyNamehint<XorOp>(rewriter, op, op.getLhs(),
3386 case ICmpPredicate::ceq:
3387 case ICmpPredicate::cne:
3388 case ICmpPredicate::weq:
3389 case ICmpPredicate::wne:
3395 if (op.getPredicate() == ICmpPredicate::eq ||
3396 op.getPredicate() == ICmpPredicate::ne) {
3401 if (!knownBits.isUnknown())
3408 if (
auto xorOp = op.getLhs().getDefiningOp<
XorOp>())
3415 if (
auto replicateOp = op.getLhs().getDefiningOp<ReplicateOp>())
3416 if (rhs.isAllOnes() || rhs.isZero()) {
3417 auto width = replicateOp.getInput().getType().getIntOrFloatBitWidth();
3420 rhs.isAllOnes() ? APInt::getAllOnes(width)
3421 : APInt::getZero(width));
3422 replaceOpWithNewOpAndCopyNamehint<ICmpOp>(
3423 rewriter, op, op.getPredicate(), replicateOp.getInput(), cst,
3433 if (Operation *opLHS = op.getLhs().getDefiningOp())
3434 if (Operation *opRHS = op.getRhs().getDefiningOp())
3435 if (isa<ConcatOp, ReplicateOp>(opLHS) &&
3436 isa<ConcatOp, ReplicateOp>(opRHS)) {
assert(baseType &&"element must be base type")
static KnownBits computeKnownBits(Value v, unsigned depth)
Given an integer SSA value, check to see if we know anything about the result of the computation.
static bool foldMuxOfUniformArrays(MuxOp op, PatternRewriter &rewriter)
static Attribute constFoldAssociativeOp(ArrayRef< Attribute > operands, hw::PEO paramOpcode)
static Attribute constFoldBinaryOp(ArrayRef< Attribute > operands, hw::PEO paramOpcode)
Performs constant folding calculate with element-wise behavior on the two attributes in operands and ...
static bool applyCmpPredicateToEqualOperands(ICmpPredicate predicate)
static ComplementMatcher< SubType > m_Complement(const SubType &subExpr)
static bool canonicalizeLogicalCstWithConcat(Operation *logicalOp, size_t concatIdx, const APInt &cst, PatternRewriter &rewriter)
When we find a logical operation (and, or, xor) with a constant e.g.
static bool narrowOperationWidth(OpTy op, bool narrowTrailingBits, PatternRewriter &rewriter)
static OpFoldResult foldDiv(Op op, ArrayRef< Attribute > constants)
static Value getCommonOperand(Op op)
Returns a single common operand that all inputs of the operation op can be traced back to,...
static bool canCombineOppositeBinCmpIntoConstant(OperandRange operands)
static void getConcatOperands(Value v, SmallVectorImpl< Value > &result)
Flatten concat and mux operands into a vector.
static Value extractOperandFromFullyAssociative(Operation *fullyAssoc, size_t operandNo, PatternRewriter &rewriter)
Given a fully associative variadic operation like (a+b+c+d), break the expression into two parts,...
static bool getMuxChainCondConstant(Value cond, Value indexValue, bool isInverted, std::function< void(hw::ConstantOp)> constantFn)
Check to see if the condition to the specified mux is an equality comparison indexValue and one or mo...
static TypedAttr getIntAttr(const APInt &value, MLIRContext *context)
static bool shouldBeFlattened(Operation *op)
Return true if the op will be flattened afterwards.
static void canonicalizeXorIcmpTrue(XorOp op, unsigned icmpOperand, PatternRewriter &rewriter)
static bool assumeMuxCondInOperand(Value muxCond, Value muxValue, bool constCond, PatternRewriter &rewriter)
If the mux condition is an operand to the op defining its true or false value, replace the condition ...
static bool extractFromReplicate(ExtractOp op, ReplicateOp replicate, PatternRewriter &rewriter)
static void combineEqualityICmpWithXorOfConstant(ICmpOp cmpOp, XorOp xorOp, const APInt &rhs, PatternRewriter &rewriter)
static size_t getTotalWidth(ArrayRef< Value > operands)
static bool foldCommonMuxOperation(MuxOp mux, Operation *trueOp, Operation *falseOp, PatternRewriter &rewriter)
This function is invoke when we find a mux with true/false operations that have the same opcode.
static std::pair< size_t, size_t > getLowestBitAndHighestBitRequired(Operation *op, bool narrowTrailingBits, size_t originalOpWidth)
static bool tryFlatteningOperands(Operation *op, PatternRewriter &rewriter)
Flattens a single input in op if hasOneUse is true and it can be defined as an Op.
static bool isOpTriviallyRecursive(Operation *op)
static LogicalResult extractConcatToConcatExtract(ExtractOp op, ConcatOp innerCat, PatternRewriter &rewriter, ArrayRef< size_t > prefixWidths={})
static bool canonicalizeIdempotentInputs(Op op, PatternRewriter &rewriter)
Canonicalize an idempotent operation op so that only one input of any kind occurs.
static bool applyCmpPredicate(ICmpPredicate predicate, const APInt &lhs, const APInt &rhs)
static void combineEqualityICmpWithKnownBitsAndConstant(ICmpOp cmpOp, const KnownBits &bitAnalysis, const APInt &rhsCst, PatternRewriter &rewriter)
Given an equality comparison with a constant value and some operand that has known bits,...
static bool hasSVAttributes(Operation *op)
static OpFoldResult foldMod(Op op, ArrayRef< Attribute > constants)
static size_t computeCommonPrefixLength(const Range &a, const Range &b)
static bool foldCommonMuxValue(MuxOp op, bool isTrueOperand, PatternRewriter &rewriter)
Fold things like mux(cond, x|y|z|a, a) -> (x|y|z)&replicate(cond)|a and mux(cond, a,...
static LogicalResult matchAndRewriteCompareConcat(ICmpOp op, Operation *lhs, Operation *rhs, PatternRewriter &rewriter)
Reduce the strength icmp(concat(...), concat(...)) by doing a element-wise comparison on common prefi...
static Value createGenericOp(Location loc, OperationName name, ArrayRef< Value > operands, OpBuilder &builder)
Create a new instance of a generic operation that only has value operands, and has a single result va...
static TypedAttr getIntAttr(MLIRContext *ctx, Type t, const APInt &value)
static std::unique_ptr< Context > context
static std::optional< APSInt > getConstant(Attribute operand)
Determine the value of a constant operand for the sake of constant folding.
create(elements, Type result_type=None)
create(array_value, low_index, ret_type)
Direction get(bool isOutput)
Returns an output direction if isOutput is true, otherwise returns an input direction.
void extractBits(OpBuilder &builder, Value val, SmallVectorImpl< Value > &bits)
Extract bits from a value.
bool foldMuxChainWithComparison(PatternRewriter &rewriter, MuxOp rootMux, bool isFalseSide, llvm::function_ref< MuxChainWithComparisonFoldingStyle(size_t indexWidth, size_t numEntries)> styleFn)
Mux chain folding that converts chains of muxes with index comparisons into array operations or balan...
Value createOrFoldNot(OpBuilder &builder, Location loc, Value value, bool twoState=false)
Create a `‘Not’' gate on a value.
MuxChainWithComparisonFoldingStyle
Enum for mux chain folding styles.
LogicalResult convertModUByPowerOfTwo(ModUOp modOp, mlir::PatternRewriter &rewriter)
KnownBits computeKnownBits(Value value)
Compute "known bits" information about the specified value - the set of bits that are guaranteed to a...
Value constructMuxTree(OpBuilder &builder, Location loc, ArrayRef< Value > selectors, ArrayRef< Value > leafNodes, Value outOfBoundsValue)
Construct a mux tree for given leaf nodes.
Value createOrFoldSExt(OpBuilder &builder, Location loc, Value value, Type destTy)
Create a sign extension operation from a value of integer type to an equal or larger integer type.
LogicalResult convertDivUByPowerOfTwo(DivUOp divOp, mlir::PatternRewriter &rewriter)
Convert unsigned division or modulo by a power of two.
uint64_t getWidth(Type t)
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.
void replaceOpAndCopyNamehint(PatternRewriter &rewriter, Operation *op, Value newValue)
A wrapper of PatternRewriter::replaceOp to propagate "sv.namehint" attribute.