20#include "mlir/IR/Matchers.h"
21#include "mlir/IR/PatternMatch.h"
22#include "llvm/ADT/APSInt.h"
23#include "llvm/ADT/STLExtras.h"
24#include "llvm/ADT/SmallPtrSet.h"
25#include "llvm/ADT/StringExtras.h"
26#include "llvm/ADT/TypeSwitch.h"
29using namespace firrtl;
33static Value
dropWrite(PatternRewriter &rewriter, OpResult old,
35 SmallPtrSet<Operation *, 8> users;
36 for (
auto *user : old.getUsers())
38 for (Operation *user : users)
39 if (
auto connect = dyn_cast<FConnectLike>(user))
40 if (connect.getDest() == old)
41 rewriter.eraseOp(user);
51 if (op->getNumRegions() != 0)
53 return mlir::isPure(op) || isa<NodeOp, WireOp>(op);
61 Operation *op = passthrough.getDefiningOp();
64 assert(op &&
"passthrough must be an operation");
65 Operation *oldOp = old.getOwner();
66 auto name = oldOp->getAttrOfType<StringAttr>(
"name");
68 op->setAttr(
"name", name);
76#include "circt/Dialect/FIRRTL/FIRRTLCanonicalization.h.inc"
84 auto resultType = type_cast<IntType>(op->getResult(0).getType());
85 if (!resultType.hasWidth())
87 for (Value operand : op->getOperands())
88 if (!type_cast<IntType>(operand.getType()).hasWidth())
95 auto t = type_dyn_cast<UIntType>(type);
96 if (!t || !t.hasWidth() || t.getWidth() != 1)
103static void updateName(PatternRewriter &rewriter, Operation *op,
108 assert((!isa<InstanceOp, RegOp, RegResetOp>(op)) &&
"Should never rename");
109 auto newName = name.getValue();
110 auto newOpName = op->getAttrOfType<StringAttr>(
"name");
113 newName =
chooseName(newOpName.getValue(), name.getValue());
115 if (!newOpName || newOpName.getValue() != newName)
116 rewriter.modifyOpInPlace(
117 op, [&] { op->setAttr(
"name", rewriter.getStringAttr(newName)); });
125 if (
auto *newOp = newValue.getDefiningOp()) {
126 auto name = op->getAttrOfType<StringAttr>(
"name");
129 rewriter.replaceOp(op, newValue);
135template <
typename OpTy,
typename... Args>
137 Operation *op, Args &&...args) {
138 auto name = op->getAttrOfType<StringAttr>(
"name");
140 rewriter.replaceOpWithNewOp<OpTy>(op, std::forward<Args>(args)...);
148 if (
auto namableOp = dyn_cast<firrtl::FNamableOp>(op))
149 return namableOp.hasDroppableName();
160static std::optional<APSInt>
162 assert(type_cast<IntType>(operand.getType()) &&
163 "getExtendedConstant is limited to integer types");
170 if (IntegerAttr result = dyn_cast_or_null<IntegerAttr>(constant))
175 if (type_cast<IntType>(operand.getType()).getWidth() == 0)
176 return APSInt(destWidth,
177 type_cast<IntType>(operand.getType()).isUnsigned());
185 if (
auto attr = dyn_cast<BoolAttr>(operand))
186 return APSInt(APInt(1, attr.getValue()));
187 if (
auto attr = dyn_cast<IntegerAttr>(operand))
188 return attr.getAPSInt();
196 return cst->isZero();
223 Operation *op, ArrayRef<Attribute> operands,
BinOpKind opKind,
224 const function_ref<APInt(
const APSInt &,
const APSInt &)> &calculate) {
225 assert(operands.size() == 2 &&
"binary op takes two operands");
228 auto resultType = type_cast<IntType>(op->getResult(0).getType());
229 if (resultType.getWidthOrSentinel() < 0)
233 if (resultType.getWidthOrSentinel() == 0)
234 return getIntAttr(resultType, APInt(0, 0, resultType.isSigned()));
240 type_cast<IntType>(op->getOperand(0).getType()).getWidthOrSentinel();
242 type_cast<IntType>(op->getOperand(1).getType()).getWidthOrSentinel();
243 if (
auto lhs = dyn_cast_or_null<IntegerAttr>(operands[0]))
244 lhsWidth = std::max<int32_t>(lhsWidth, lhs.getValue().getBitWidth());
245 if (
auto rhs = dyn_cast_or_null<IntegerAttr>(operands[1]))
246 rhsWidth = std::max<int32_t>(rhsWidth, rhs.getValue().getBitWidth());
250 int32_t operandWidth;
253 operandWidth = resultType.getWidthOrSentinel();
258 operandWidth = std::max(1, std::max(lhsWidth, rhsWidth));
262 std::max(std::max(lhsWidth, rhsWidth), resultType.getWidthOrSentinel());
273 APInt resultValue = calculate(*lhs, *rhs);
278 resultValue = resultValue.trunc(resultType.getWidthOrSentinel());
280 assert((
unsigned)resultType.getWidthOrSentinel() ==
281 resultValue.getBitWidth());
294 Operation *op, PatternRewriter &rewriter,
295 const function_ref<OpFoldResult(ArrayRef<Attribute>)> &canonicalize) {
297 if (op->getNumResults() != 1)
299 auto type = type_dyn_cast<FIRRTLBaseType>(op->getResult(0).getType());
304 auto width = type.getBitWidthOrSentinel();
309 SmallVector<Attribute, 3> constOperands;
310 constOperands.reserve(op->getNumOperands());
311 for (
auto operand : op->getOperands()) {
313 if (
auto *defOp = operand.getDefiningOp())
314 TypeSwitch<Operation *>(defOp).Case<ConstantOp, SpecialConstantOp>(
315 [&](
auto op) { attr = op.getValueAttr(); });
316 constOperands.push_back(attr);
321 auto result = canonicalize(constOperands);
325 if (
auto cst = dyn_cast<Attribute>(result))
326 resultValue = op->getDialect()
327 ->materializeConstant(rewriter, cst, type, op->getLoc())
330 resultValue = cast<Value>(result);
334 type_cast<FIRRTLBaseType>(resultValue.getType()).getBitWidthOrSentinel())
335 resultValue = PadPrimOp::create(rewriter, op->getLoc(), resultValue, width);
338 if (type_isa<SIntType>(type) && type_isa<UIntType>(resultValue.getType()))
339 resultValue = AsSIntPrimOp::create(rewriter, op->getLoc(), resultValue);
340 else if (type_isa<UIntType>(type) &&
341 type_isa<SIntType>(resultValue.getType()))
342 resultValue = AsUIntPrimOp::create(rewriter, op->getLoc(), resultValue);
344 assert(type == resultValue.getType() &&
"canonicalization changed type");
352 return bitWidth > 0 ? APInt::getMaxValue(bitWidth) : APInt();
358 return bitWidth > 0 ? APInt::getSignedMinValue(bitWidth) : APInt();
364 return bitWidth > 0 ? APInt::getSignedMaxValue(bitWidth) : APInt();
371OpFoldResult ConstantOp::fold(FoldAdaptor adaptor) {
372 assert(adaptor.getOperands().empty() &&
"constant has no operands");
373 return getValueAttr();
376OpFoldResult SpecialConstantOp::fold(FoldAdaptor adaptor) {
377 assert(adaptor.getOperands().empty() &&
"constant has no operands");
378 return getValueAttr();
381OpFoldResult AggregateConstantOp::fold(FoldAdaptor adaptor) {
382 assert(adaptor.getOperands().empty() &&
"constant has no operands");
383 return getFieldsAttr();
386OpFoldResult StringConstantOp::fold(FoldAdaptor adaptor) {
387 assert(adaptor.getOperands().empty() &&
"constant has no operands");
388 return getValueAttr();
391OpFoldResult FIntegerConstantOp::fold(FoldAdaptor adaptor) {
392 assert(adaptor.getOperands().empty() &&
"constant has no operands");
393 return getValueAttr();
396OpFoldResult BoolConstantOp::fold(FoldAdaptor adaptor) {
397 assert(adaptor.getOperands().empty() &&
"constant has no operands");
398 return getValueAttr();
401OpFoldResult DoubleConstantOp::fold(FoldAdaptor adaptor) {
402 assert(adaptor.getOperands().empty() &&
"constant has no operands");
403 return getValueAttr();
410OpFoldResult AddPrimOp::fold(FoldAdaptor adaptor) {
413 [=](
const APSInt &a,
const APSInt &b) { return a + b; });
416void AddPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
418 results.insert<patterns::moveConstAdd, patterns::AddOfZero,
419 patterns::AddOfSelf, patterns::AddOfPad>(
context);
422OpFoldResult SubPrimOp::fold(FoldAdaptor adaptor) {
425 [=](
const APSInt &a,
const APSInt &b) { return a - b; });
428void SubPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
430 results.insert<patterns::SubOfZero, patterns::SubFromZeroSigned,
431 patterns::SubFromZeroUnsigned, patterns::SubOfSelf,
432 patterns::SubOfPadL, patterns::SubOfPadR>(
context);
435OpFoldResult MulPrimOp::fold(FoldAdaptor adaptor) {
447 [=](
const APSInt &a,
const APSInt &b) { return a * b; });
450OpFoldResult DivPrimOp::fold(FoldAdaptor adaptor) {
457 if (getLhs() == getRhs()) {
458 auto width = getType().base().getWidthOrSentinel();
463 return getIntAttr(getType(), APInt(width, 1));
480 if (
auto rhsCst = dyn_cast_or_null<IntegerAttr>(adaptor.getRhs()))
481 if (rhsCst.getValue().isOne() && getLhs().getType() == getType())
486 [=](
const APSInt &a,
const APSInt &b) -> APInt {
489 return APInt(a.getBitWidth(), 0);
493OpFoldResult RemPrimOp::fold(FoldAdaptor adaptor) {
500 if (getLhs() == getRhs())
514 [=](
const APSInt &a,
const APSInt &b) -> APInt {
517 return APInt(a.getBitWidth(), 0);
521OpFoldResult DShlPrimOp::fold(FoldAdaptor adaptor) {
524 [=](
const APSInt &a,
const APSInt &b) -> APInt { return a.shl(b); });
527OpFoldResult DShlwPrimOp::fold(FoldAdaptor adaptor) {
530 [=](
const APSInt &a,
const APSInt &b) -> APInt { return a.shl(b); });
533OpFoldResult DShrPrimOp::fold(FoldAdaptor adaptor) {
536 [=](
const APSInt &a,
const APSInt &b) -> APInt {
537 return getType().base().isUnsigned() || !a.getBitWidth() ? a.lshr(b)
543OpFoldResult AndPrimOp::fold(FoldAdaptor adaptor) {
546 if (rhsCst->isZero())
550 if (rhsCst->isAllOnes() && getLhs().getType() == getType() &&
551 getRhs().getType() == getType())
557 if (lhsCst->isZero())
561 if (lhsCst->isAllOnes() && getLhs().getType() == getType() &&
562 getRhs().getType() == getType())
567 if (getLhs() == getRhs() && getRhs().getType() == getType())
572 [](
const APSInt &a,
const APSInt &b) -> APInt { return a & b; });
575void AndPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
578 .insert<patterns::extendAnd, patterns::moveConstAnd, patterns::AndOfZero,
579 patterns::AndOfAllOne, patterns::AndOfSelf, patterns::AndOfPad,
580 patterns::AndOfAsSIntL, patterns::AndOfAsSIntR>(
context);
583OpFoldResult OrPrimOp::fold(FoldAdaptor adaptor) {
586 if (rhsCst->isZero() && getLhs().getType() == getType())
590 if (rhsCst->isAllOnes() && getRhs().getType() == getType() &&
591 getLhs().getType() == getType())
597 if (lhsCst->isZero() && getRhs().getType() == getType())
601 if (lhsCst->isAllOnes() && getLhs().getType() == getType() &&
602 getRhs().getType() == getType())
607 if (getLhs() == getRhs() && getRhs().getType() == getType())
612 [](
const APSInt &a,
const APSInt &b) -> APInt { return a | b; });
615void OrPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
617 results.insert<patterns::extendOr, patterns::moveConstOr, patterns::OrOfZero,
618 patterns::OrOfAllOne, patterns::OrOfSelf, patterns::OrOfPad,
622OpFoldResult XorPrimOp::fold(FoldAdaptor adaptor) {
625 if (rhsCst->isZero() &&
631 if (lhsCst->isZero() &&
636 if (getLhs() == getRhs())
639 APInt(std::max(getType().base().getWidthOrSentinel(), 0), 0));
643 [](
const APSInt &a,
const APSInt &b) -> APInt { return a ^ b; });
646void XorPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
648 results.insert<patterns::extendXor, patterns::moveConstXor,
649 patterns::XorOfZero, patterns::XorOfSelf, patterns::XorOfPad>(
653void LEQPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
655 results.insert<patterns::LEQWithConstLHS>(
context);
658OpFoldResult LEQPrimOp::fold(FoldAdaptor adaptor) {
659 bool isUnsigned = getLhs().getType().base().isUnsigned();
662 if (getLhs() == getRhs())
666 if (
auto width = getLhs().getType().base().
getWidth()) {
668 auto commonWidth = std::max<int32_t>(*width, rhsCst->getBitWidth());
669 commonWidth = std::max(commonWidth, 1);
680 if (isUnsigned && rhsCst->zext(commonWidth)
693 [=](
const APSInt &a,
const APSInt &b) -> APInt {
694 return APInt(1, a <= b);
698void LTPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
700 results.insert<patterns::LTWithConstLHS>(
context);
703OpFoldResult LTPrimOp::fold(FoldAdaptor adaptor) {
704 IntType lhsType = getLhs().getType();
708 if (getLhs() == getRhs())
718 if (
auto width = lhsType.
getWidth()) {
720 auto commonWidth = std::max<int32_t>(*width, rhsCst->getBitWidth());
721 commonWidth = std::max(commonWidth, 1);
732 if (isUnsigned && rhsCst->zext(commonWidth)
745 [=](
const APSInt &a,
const APSInt &b) -> APInt {
746 return APInt(1, a < b);
750void GEQPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
752 results.insert<patterns::GEQWithConstLHS>(
context);
755OpFoldResult GEQPrimOp::fold(FoldAdaptor adaptor) {
756 IntType lhsType = getLhs().getType();
760 if (getLhs() == getRhs())
765 if (rhsCst->isZero() && isUnsigned)
770 if (
auto width = lhsType.
getWidth()) {
772 auto commonWidth = std::max<int32_t>(*width, rhsCst->getBitWidth());
773 commonWidth = std::max(commonWidth, 1);
776 if (isUnsigned && rhsCst->zext(commonWidth)
797 [=](
const APSInt &a,
const APSInt &b) -> APInt {
798 return APInt(1, a >= b);
802void GTPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
804 results.insert<patterns::GTWithConstLHS>(
context);
807OpFoldResult GTPrimOp::fold(FoldAdaptor adaptor) {
808 IntType lhsType = getLhs().getType();
812 if (getLhs() == getRhs())
816 if (
auto width = lhsType.
getWidth()) {
818 auto commonWidth = std::max<int32_t>(*width, rhsCst->getBitWidth());
819 commonWidth = std::max(commonWidth, 1);
822 if (isUnsigned && rhsCst->zext(commonWidth)
843 [=](
const APSInt &a,
const APSInt &b) -> APInt {
844 return APInt(1, a > b);
848OpFoldResult EQPrimOp::fold(FoldAdaptor adaptor) {
850 if (getLhs() == getRhs())
856 if (rhsCst->isAllOnes() && getLhs().getType() == getType() &&
857 getRhs().getType() == getType())
863 [=](
const APSInt &a,
const APSInt &b) -> APInt {
864 return APInt(1, a == b);
868LogicalResult EQPrimOp::canonicalize(EQPrimOp op, PatternRewriter &rewriter) {
870 op, rewriter, [&](ArrayRef<Attribute> operands) -> OpFoldResult {
872 auto width = op.getLhs().getType().getBitWidthOrSentinel();
875 if (rhsCst->isZero() && op.getLhs().getType() == op.getType() &&
876 op.getRhs().getType() == op.getType()) {
877 return NotPrimOp::create(rewriter, op.getLoc(), op.getLhs())
882 if (rhsCst->isZero() && width > 1) {
883 auto orrOp = OrRPrimOp::create(rewriter, op.getLoc(), op.getLhs());
884 return NotPrimOp::create(rewriter, op.getLoc(), orrOp).getResult();
888 if (rhsCst->isAllOnes() && width > 1 &&
889 op.getLhs().getType() == op.getRhs().getType()) {
890 return AndRPrimOp::create(rewriter, op.getLoc(), op.getLhs())
898OpFoldResult NEQPrimOp::fold(FoldAdaptor adaptor) {
900 if (getLhs() == getRhs())
906 if (rhsCst->isZero() && getLhs().getType() == getType() &&
907 getRhs().getType() == getType())
913 [=](
const APSInt &a,
const APSInt &b) -> APInt {
914 return APInt(1, a != b);
918LogicalResult NEQPrimOp::canonicalize(NEQPrimOp op, PatternRewriter &rewriter) {
920 op, rewriter, [&](ArrayRef<Attribute> operands) -> OpFoldResult {
922 auto width = op.getLhs().getType().getBitWidthOrSentinel();
925 if (rhsCst->isAllOnes() && op.getLhs().getType() == op.getType() &&
926 op.getRhs().getType() == op.getType()) {
927 return NotPrimOp::create(rewriter, op.getLoc(), op.getLhs())
932 if (rhsCst->isZero() && width > 1) {
933 return OrRPrimOp::create(rewriter, op.getLoc(), op.getLhs())
938 if (rhsCst->isAllOnes() && width > 1 &&
939 op.getLhs().getType() == op.getRhs().getType()) {
941 AndRPrimOp::create(rewriter, op.getLoc(), op.getLhs());
942 return NotPrimOp::create(rewriter, op.getLoc(), andrOp).getResult();
950OpFoldResult IntegerAddOp::fold(FoldAdaptor adaptor) {
956OpFoldResult IntegerMulOp::fold(FoldAdaptor adaptor) {
962OpFoldResult IntegerShrOp::fold(FoldAdaptor adaptor) {
966 return IntegerAttr::get(IntegerType::get(getContext(),
967 lhsCst->getBitWidth(),
968 IntegerType::Signed),
969 lhsCst->ashr(*rhsCst));
972 if (rhsCst->isZero())
979OpFoldResult IntegerShlOp::fold(FoldAdaptor adaptor) {
984 return IntegerAttr::get(IntegerType::get(getContext(),
985 lhsCst->getBitWidth(),
986 IntegerType::Signed),
987 lhsCst->shl(*rhsCst));
990 if (rhsCst->isZero())
1001OpFoldResult SizeOfIntrinsicOp::fold(FoldAdaptor) {
1002 auto base = getInput().getType();
1009OpFoldResult IsXIntrinsicOp::fold(FoldAdaptor adaptor) {
1016OpFoldResult AsSIntPrimOp::fold(FoldAdaptor adaptor) {
1024 if (getType().base().hasWidth())
1031void AsSIntPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1033 results.insert<patterns::StoUtoS>(
context);
1036OpFoldResult AsUIntPrimOp::fold(FoldAdaptor adaptor) {
1044 if (getType().base().hasWidth())
1051void AsUIntPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1053 results.insert<patterns::UtoStoU>(
context);
1056OpFoldResult AsAsyncResetPrimOp::fold(FoldAdaptor adaptor) {
1058 if (getInput().getType() == getType())
1063 return BoolAttr::get(getContext(), cst->getBoolValue());
1068OpFoldResult AsResetPrimOp::fold(FoldAdaptor adaptor) {
1070 return BoolAttr::get(getContext(), cst->getBoolValue());
1074OpFoldResult AsClockPrimOp::fold(FoldAdaptor adaptor) {
1076 if (getInput().getType() == getType())
1081 return BoolAttr::get(getContext(), cst->getBoolValue());
1086OpFoldResult CvtPrimOp::fold(FoldAdaptor adaptor) {
1092 getType().base().getWidthOrSentinel()))
1098void CvtPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1100 results.insert<patterns::CVTSigned, patterns::CVTUnSigned>(
context);
1103OpFoldResult NegPrimOp::fold(FoldAdaptor adaptor) {
1110 getType().base().getWidthOrSentinel()))
1111 return getIntAttr(getType(), APInt((*cst).getBitWidth(), 0) - *cst);
1116OpFoldResult NotPrimOp::fold(FoldAdaptor adaptor) {
1121 getType().base().getWidthOrSentinel()))
1127void NotPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1129 results.insert<patterns::NotNot, patterns::NotEq, patterns::NotNeq,
1130 patterns::NotLeq, patterns::NotLt, patterns::NotGeq,
1137 : RewritePattern(opName, 0,
context) {}
1143 ConstantOp constantOp,
1144 SmallVectorImpl<Value> &remaining)
const = 0;
1151 mlir::PatternRewriter &rewriter)
const override {
1153 auto catOp = op->getOperand(0).getDefiningOp<CatPrimOp>();
1157 SmallVector<Value> nonConstantOperands;
1160 for (
auto operand : catOp.getInputs()) {
1161 if (
auto constantOp = operand.getDefiningOp<ConstantOp>()) {
1163 if (
handleConstant(rewriter, op, constantOp, nonConstantOperands))
1167 nonConstantOperands.push_back(operand);
1172 if (nonConstantOperands.empty()) {
1173 replaceOpWithNewOpAndCopyName<ConstantOp>(
1174 rewriter, op, cast<IntType>(op->getResult(0).getType()),
1180 if (nonConstantOperands.size() == 1) {
1181 rewriter.modifyOpInPlace(
1182 op, [&] { op->setOperand(0, nonConstantOperands.front()); });
1187 if (catOp->hasOneUse() &&
1188 nonConstantOperands.size() < catOp->getNumOperands()) {
1189 replaceOpWithNewOpAndCopyName<CatPrimOp>(rewriter, catOp,
1190 nonConstantOperands);
1203 SmallVectorImpl<Value> &remaining)
const override {
1204 if (value.getValue().isZero())
1207 replaceOpWithNewOpAndCopyName<ConstantOp>(
1208 rewriter, op, cast<IntType>(op->getResult(0).getType()),
1221 SmallVectorImpl<Value> &remaining)
const override {
1222 if (value.getValue().isAllOnes())
1225 replaceOpWithNewOpAndCopyName<ConstantOp>(
1226 rewriter, op, cast<IntType>(op->getResult(0).getType()),
1239 SmallVectorImpl<Value> &remaining)
const override {
1240 if (value.getValue().isZero())
1242 remaining.push_back(value);
1248OpFoldResult AndRPrimOp::fold(FoldAdaptor adaptor) {
1252 if (getInput().getType().getBitWidthOrSentinel() == 0)
1257 return getIntAttr(getType(), APInt(1, cst->isAllOnes()));
1261 if (
isUInt1(getInput().getType()))
1267void AndRPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1269 results.insert<patterns::AndRasSInt, patterns::AndRasUInt, patterns::AndRPadU,
1270 patterns::AndRPadS, patterns::AndRCatAndR_left,
1274OpFoldResult OrRPrimOp::fold(FoldAdaptor adaptor) {
1278 if (getInput().getType().getBitWidthOrSentinel() == 0)
1283 return getIntAttr(getType(), APInt(1, !cst->isZero()));
1287 if (
isUInt1(getInput().getType()))
1293void OrRPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1295 results.insert<patterns::OrRasSInt, patterns::OrRasUInt, patterns::OrRPadU,
1296 patterns::OrRCatOrR_left, patterns::OrRCatOrR_right,
OrRCat>(
1300OpFoldResult XorRPrimOp::fold(FoldAdaptor adaptor) {
1304 if (getInput().getType().getBitWidthOrSentinel() == 0)
1309 return getIntAttr(getType(), APInt(1, cst->popcount() & 1));
1312 if (
isUInt1(getInput().getType()))
1318void XorRPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1321 .insert<patterns::XorRasSInt, patterns::XorRasUInt, patterns::XorRPadU,
1322 patterns::XorRCatXorR_left, patterns::XorRCatXorR_right,
XorRCat>(
1330OpFoldResult CatPrimOp::fold(FoldAdaptor adaptor) {
1331 auto inputs = getInputs();
1332 auto inputAdaptors = adaptor.getInputs();
1339 if (inputs.size() == 1 && inputs[0].getType() == getType())
1347 SmallVector<Value> nonZeroInputs;
1348 SmallVector<Attribute> nonZeroAttributes;
1349 bool allConstant =
true;
1350 for (
auto [input, attr] :
llvm::zip(inputs, inputAdaptors)) {
1351 auto inputType = type_cast<IntType>(input.getType());
1352 if (inputType.getBitWidthOrSentinel() != 0) {
1353 nonZeroInputs.push_back(input);
1355 allConstant =
false;
1356 if (nonZeroInputs.size() > 1 && !allConstant)
1362 if (nonZeroInputs.empty())
1366 if (nonZeroInputs.size() == 1 && nonZeroInputs[0].getType() == getType())
1367 return nonZeroInputs[0];
1373 SmallVector<APInt> constants;
1374 for (
auto inputAdaptor : inputAdaptors) {
1376 constants.push_back(*cst);
1381 assert(!constants.empty());
1383 APInt result = constants[0];
1384 for (
size_t i = 1; i < constants.size(); ++i)
1385 result = result.concat(constants[i]);
1390void DShlPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1392 results.insert<patterns::DShlOfConstant>(
context);
1395void DShrPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1397 results.insert<patterns::DShrOfConstant>(
context);
1405 using OpRewritePattern::OpRewritePattern;
1408 matchAndRewrite(CatPrimOp cat,
1409 mlir::PatternRewriter &rewriter)
const override {
1411 cat.getType().getBitWidthOrSentinel() == 0)
1415 if (cat->hasOneUse() && isa<CatPrimOp>(*cat->getUsers().begin()))
1419 SmallVector<Value> operands;
1420 SmallVector<Value> worklist;
1421 auto pushOperands = [&worklist](CatPrimOp op) {
1422 for (
auto operand :
llvm::reverse(op.getInputs()))
1423 worklist.push_back(operand);
1426 bool hasSigned =
false, hasUnsigned =
false;
1427 while (!worklist.empty()) {
1428 auto value = worklist.pop_back_val();
1429 auto catOp = value.getDefiningOp<CatPrimOp>();
1431 operands.push_back(value);
1432 (type_isa<UIntType>(value.getType()) ? hasUnsigned : hasSigned) =
true;
1436 pushOperands(catOp);
1441 auto castToUIntIfSigned = [&](Value value) -> Value {
1442 if (type_isa<UIntType>(value.getType()))
1444 return AsUIntPrimOp::create(rewriter, value.getLoc(), value);
1447 assert(operands.size() >= 1 &&
"zero width cast must be rejected");
1449 if (operands.size() == 1) {
1450 rewriter.replaceOp(cat, castToUIntIfSigned(operands[0]));
1454 if (operands.size() == cat->getNumOperands())
1458 if (hasSigned && hasUnsigned)
1459 for (
auto &operand : operands)
1460 operand = castToUIntIfSigned(operand);
1462 replaceOpWithNewOpAndCopyName<CatPrimOp>(rewriter, cat, cat.getType(),
1471 using OpRewritePattern::OpRewritePattern;
1474 matchAndRewrite(CatPrimOp cat,
1475 mlir::PatternRewriter &rewriter)
const override {
1479 SmallVector<Value> operands;
1481 for (
size_t i = 0; i < cat->getNumOperands(); ++i) {
1482 auto cst = cat.getInputs()[i].getDefiningOp<ConstantOp>();
1484 operands.push_back(cat.getInputs()[i]);
1487 APSInt value = cst.getValue();
1489 for (; j < cat->getNumOperands(); ++j) {
1490 auto nextCst = cat.getInputs()[j].getDefiningOp<ConstantOp>();
1493 value = value.concat(nextCst.getValue());
1498 operands.push_back(cst);
1501 operands.push_back(ConstantOp::create(rewriter, cat.getLoc(), value));
1507 if (operands.size() == cat->getNumOperands())
1510 replaceOpWithNewOpAndCopyName<CatPrimOp>(rewriter, cat, cat.getType(),
1519void CatPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1521 results.insert<patterns::CatBitsBits, patterns::CatDoubleConst,
1522 patterns::CatCast, FlattenCat, CatOfConstant>(
context);
1529OpFoldResult StringConcatOp::fold(FoldAdaptor adaptor) {
1531 if (getInputs().size() == 1)
1532 return getInputs()[0];
1535 if (!llvm::all_of(adaptor.getInputs(), [](Attribute operand) {
1536 return isa_and_nonnull<StringAttr>(operand);
1541 SmallString<64> result;
1542 for (
auto operand : adaptor.getInputs())
1543 result += cast<StringAttr>(operand).getValue();
1545 return StringAttr::get(getContext(), result);
1553 using OpRewritePattern::OpRewritePattern;
1556 matchAndRewrite(StringConcatOp concat,
1557 mlir::PatternRewriter &rewriter)
const override {
1561 bool hasNestedConcat = llvm::any_of(concat.getInputs(), [](Value operand) {
1562 auto nestedConcat = operand.getDefiningOp<StringConcatOp>();
1563 return nestedConcat && operand.hasOneUse();
1566 if (!hasNestedConcat)
1570 SmallVector<Value> flatOperands;
1571 for (
auto input : concat.getInputs()) {
1572 if (
auto nestedConcat = input.getDefiningOp<StringConcatOp>();
1573 nestedConcat && input.hasOneUse())
1574 llvm::append_range(flatOperands, nestedConcat.getInputs());
1576 flatOperands.push_back(input);
1579 rewriter.modifyOpInPlace(concat,
1580 [&]() { concat->setOperands(flatOperands); });
1587class MergeAdjacentStringConstants
1590 using OpRewritePattern::OpRewritePattern;
1593 matchAndRewrite(StringConcatOp concat,
1594 mlir::PatternRewriter &rewriter)
const override {
1596 SmallVector<Value> newOperands;
1597 SmallString<64> accumulatedLit;
1598 SmallVector<StringConstantOp> accumulatedOps;
1599 bool changed =
false;
1601 auto flushLiterals = [&]() {
1602 if (accumulatedOps.empty())
1606 if (accumulatedOps.size() == 1) {
1607 newOperands.push_back(accumulatedOps[0]);
1610 auto newLit = rewriter.createOrFold<StringConstantOp>(
1611 concat.getLoc(), StringAttr::get(getContext(), accumulatedLit));
1612 newOperands.push_back(newLit);
1615 accumulatedLit.clear();
1616 accumulatedOps.clear();
1619 for (
auto operand : concat.getInputs()) {
1620 if (
auto litOp = operand.getDefiningOp<StringConstantOp>()) {
1622 if (litOp.getValue().empty()) {
1626 accumulatedLit += litOp.getValue();
1627 accumulatedOps.push_back(litOp);
1630 newOperands.push_back(operand);
1641 if (newOperands.empty())
1642 return rewriter.replaceOpWithNewOp<StringConstantOp>(
1643 concat, StringAttr::get(getContext(),
"")),
1647 rewriter.modifyOpInPlace(concat,
1648 [&]() { concat->setOperands(newOperands); });
1655void StringConcatOp::getCanonicalizationPatterns(RewritePatternSet &results,
1657 results.insert<FlattenStringConcat, MergeAdjacentStringConstants>(
context);
1664OpFoldResult PropEqOp::fold(FoldAdaptor adaptor) {
1665 auto lhsAttr = adaptor.getLhs();
1666 auto rhsAttr = adaptor.getRhs();
1667 if (!lhsAttr || !rhsAttr)
1670 return BoolAttr::get(getContext(), lhsAttr == rhsAttr);
1679 if (
auto boolAttr = dyn_cast_or_null<BoolAttr>(attr))
1680 return boolAttr.getValue();
1681 return std::nullopt;
1684OpFoldResult BoolAndOp::fold(FoldAdaptor adaptor) {
1688 return BoolAttr::get(getContext(), *lhs && *rhs);
1690 if ((lhs && !*lhs) || (rhs && !*rhs))
1691 return BoolAttr::get(getContext(),
false);
1700OpFoldResult BoolOrOp::fold(FoldAdaptor adaptor) {
1704 return BoolAttr::get(getContext(), *lhs || *rhs);
1706 if ((lhs && *lhs) || (rhs && *rhs))
1707 return BoolAttr::get(getContext(),
true);
1716OpFoldResult BoolXorOp::fold(FoldAdaptor adaptor) {
1720 return BoolAttr::get(getContext(), *lhs ^ *rhs);
1729OpFoldResult BitCastOp::fold(FoldAdaptor adaptor) {
1732 if (op.getType() == op.getInput().getType())
1733 return op.getInput();
1737 if (BitCastOp in = dyn_cast_or_null<BitCastOp>(op.getInput().getDefiningOp()))
1738 if (op.getType() == in.getInput().getType())
1739 return in.getInput();
1744OpFoldResult BitsPrimOp::fold(FoldAdaptor adaptor) {
1745 IntType inputType = getInput().getType();
1746 IntType resultType = getType();
1748 if (inputType == getType() && resultType.
hasWidth())
1755 cst->extractBits(getHi() - getLo() + 1, getLo()));
1761 using OpRewritePattern::OpRewritePattern;
1765 mlir::PatternRewriter &rewriter)
const override {
1766 auto cat = bits.getInput().getDefiningOp<CatPrimOp>();
1769 int32_t bitPos = bits.getLo();
1770 auto resultWidth = type_cast<UIntType>(bits.getType()).getWidthOrSentinel();
1771 if (resultWidth < 0)
1773 for (
auto operand : llvm::reverse(cat.getInputs())) {
1775 type_cast<IntType>(operand.getType()).getWidthOrSentinel();
1776 if (operandWidth < 0)
1778 if (bitPos < operandWidth) {
1779 if (bitPos + resultWidth <= operandWidth) {
1780 auto newBits = rewriter.createOrFold<BitsPrimOp>(
1781 bits.getLoc(), operand, bitPos + resultWidth - 1, bitPos);
1787 bitPos -= operandWidth;
1793void BitsPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1796 .insert<patterns::BitsOfBits, patterns::BitsOfMux, patterns::BitsOfAsUInt,
1804 unsigned loBit, PatternRewriter &rewriter) {
1805 auto resType = type_cast<IntType>(op->getResult(0).getType());
1806 if (type_cast<IntType>(value.getType()).getWidth() != resType.getWidth())
1807 value = BitsPrimOp::create(rewriter, op->getLoc(), value, hiBit, loBit);
1809 if (resType.isSigned() && !type_cast<IntType>(value.getType()).isSigned()) {
1810 value = rewriter.createOrFold<AsSIntPrimOp>(op->getLoc(), resType, value);
1811 }
else if (resType.isUnsigned() &&
1812 !type_cast<IntType>(value.getType()).isUnsigned()) {
1813 value = rewriter.createOrFold<AsUIntPrimOp>(op->getLoc(), resType, value);
1815 rewriter.replaceOp(op, value);
1818template <
typename OpTy>
1819static OpFoldResult
foldMux(OpTy op,
typename OpTy::FoldAdaptor adaptor) {
1821 if (op.getType().getBitWidthOrSentinel() == 0)
1823 APInt(0, 0, op.getType().isSignedInteger()));
1826 if (op.getHigh() == op.getLow() && op.getHigh().getType() == op.getType())
1827 return op.getHigh();
1832 if (op.getType().getBitWidthOrSentinel() < 0)
1837 if (cond->isZero() && op.getLow().getType() == op.getType())
1839 if (!cond->isZero() && op.getHigh().getType() == op.getType())
1840 return op.getHigh();
1844 if (
auto lowCst =
getConstant(adaptor.getLow())) {
1846 if (
auto highCst =
getConstant(adaptor.getHigh())) {
1848 if (highCst->getBitWidth() == lowCst->getBitWidth() &&
1849 *highCst == *lowCst)
1852 if (
auto intType = type_dyn_cast<IntType>(op.getType()))
1853 if (intType.hasWidth() &&
1854 (
unsigned)intType.getWidthOrSentinel() == highCst->getBitWidth())
1857 if (highCst->isOne() && lowCst->isZero() &&
1858 op.getType() == op.getSel().getType())
1871OpFoldResult MuxPrimOp::fold(FoldAdaptor adaptor) {
1872 return foldMux(*
this, adaptor);
1875OpFoldResult Mux2CellIntrinsicOp::fold(FoldAdaptor adaptor) {
1876 return foldMux(*
this, adaptor);
1879OpFoldResult Mux4CellIntrinsicOp::fold(FoldAdaptor adaptor) {
return {}; }
1888 using OpRewritePattern::OpRewritePattern;
1891 matchAndRewrite(MuxPrimOp mux,
1892 mlir::PatternRewriter &rewriter)
const override {
1893 auto width = mux.getType().getBitWidthOrSentinel();
1897 auto pad = [&](Value input) -> Value {
1899 type_cast<FIRRTLBaseType>(input.getType()).getBitWidthOrSentinel();
1900 if (inputWidth < 0 || width == inputWidth)
1902 return PadPrimOp::create(rewriter, mux.getLoc(), mux.getType(), input,
1907 auto newHigh = pad(mux.getHigh());
1908 auto newLow = pad(mux.getLow());
1909 if (newHigh == mux.getHigh() && newLow == mux.getLow())
1912 replaceOpWithNewOpAndCopyName<MuxPrimOp>(
1913 rewriter, mux, mux.getType(), ValueRange{mux.getSel(), newHigh, newLow},
1923 using OpRewritePattern::OpRewritePattern;
1925 static const int depthLimit = 5;
1927 Value updateOrClone(MuxPrimOp mux, Value high, Value low,
1928 mlir::PatternRewriter &rewriter,
1929 bool updateInPlace)
const {
1930 if (updateInPlace) {
1931 rewriter.modifyOpInPlace(mux, [&] {
1932 mux.setOperand(1, high);
1933 mux.setOperand(2, low);
1937 rewriter.setInsertionPointAfter(mux);
1938 return MuxPrimOp::create(rewriter, mux.getLoc(), mux.getType(),
1939 ValueRange{mux.getSel(), high, low})
1944 Value tryCondTrue(Value op, Value cond, mlir::PatternRewriter &rewriter,
1945 bool updateInPlace,
int limit)
const {
1946 MuxPrimOp mux = op.getDefiningOp<MuxPrimOp>();
1949 if (mux.getSel() == cond)
1950 return mux.getHigh();
1951 if (limit > depthLimit)
1953 updateInPlace &= mux->hasOneUse();
1955 if (Value v = tryCondTrue(mux.getHigh(), cond, rewriter, updateInPlace,
1957 return updateOrClone(mux, v, mux.getLow(), rewriter, updateInPlace);
1960 tryCondTrue(mux.getLow(), cond, rewriter, updateInPlace, limit + 1))
1961 return updateOrClone(mux, mux.getHigh(), v, rewriter, updateInPlace);
1966 Value tryCondFalse(Value op, Value cond, mlir::PatternRewriter &rewriter,
1967 bool updateInPlace,
int limit)
const {
1968 MuxPrimOp mux = op.getDefiningOp<MuxPrimOp>();
1971 if (mux.getSel() == cond)
1972 return mux.getLow();
1973 if (limit > depthLimit)
1975 updateInPlace &= mux->hasOneUse();
1977 if (Value v = tryCondFalse(mux.getHigh(), cond, rewriter, updateInPlace,
1979 return updateOrClone(mux, v, mux.getLow(), rewriter, updateInPlace);
1981 if (Value v = tryCondFalse(mux.getLow(), cond, rewriter, updateInPlace,
1983 return updateOrClone(mux, mux.getHigh(), v, rewriter, updateInPlace);
1989 matchAndRewrite(MuxPrimOp mux,
1990 mlir::PatternRewriter &rewriter)
const override {
1991 auto width = mux.getType().getBitWidthOrSentinel();
1995 if (Value v = tryCondTrue(mux.getHigh(), mux.getSel(), rewriter,
true, 0)) {
1996 rewriter.modifyOpInPlace(mux, [&] { mux.setOperand(1, v); });
2000 if (Value v = tryCondFalse(mux.getLow(), mux.getSel(), rewriter,
true, 0)) {
2001 rewriter.modifyOpInPlace(mux, [&] { mux.setOperand(2, v); });
2010void MuxPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
2013 .add<MuxPad, MuxSharedCond, patterns::MuxEQOperands,
2014 patterns::MuxEQOperandsSwapped, patterns::MuxNEQ, patterns::MuxNot,
2015 patterns::MuxSameTrue, patterns::MuxSameFalse,
2016 patterns::NarrowMuxLHS, patterns::NarrowMuxRHS, patterns::MuxPadSel>(
2020void Mux2CellIntrinsicOp::getCanonicalizationPatterns(
2021 RewritePatternSet &results, MLIRContext *
context) {
2022 results.add<patterns::Mux2PadSel>(
context);
2025void Mux4CellIntrinsicOp::getCanonicalizationPatterns(
2026 RewritePatternSet &results, MLIRContext *
context) {
2027 results.add<patterns::Mux4PadSel>(
context);
2030OpFoldResult PadPrimOp::fold(FoldAdaptor adaptor) {
2031 auto input = this->getInput();
2034 if (input.getType() == getType())
2038 auto inputType = input.getType().base();
2045 auto destWidth = getType().base().getWidthOrSentinel();
2046 if (destWidth == -1)
2049 if (inputType.
isSigned() && cst->getBitWidth())
2050 return getIntAttr(getType(), cst->sext(destWidth));
2051 return getIntAttr(getType(), cst->zext(destWidth));
2057OpFoldResult ShlPrimOp::fold(FoldAdaptor adaptor) {
2058 auto input = this->getInput();
2059 IntType inputType = input.getType();
2060 int shiftAmount = getAmount();
2063 if (shiftAmount == 0)
2069 if (inputWidth != -1) {
2070 auto resultWidth = inputWidth + shiftAmount;
2071 shiftAmount = std::min(shiftAmount, resultWidth);
2072 return getIntAttr(getType(), cst->zext(resultWidth).shl(shiftAmount));
2078OpFoldResult ShrPrimOp::fold(FoldAdaptor adaptor) {
2079 auto input = this->getInput();
2080 IntType inputType = input.getType();
2081 int shiftAmount = getAmount();
2087 if (shiftAmount == 0 && inputWidth > 0)
2090 if (inputWidth == -1)
2092 if (inputWidth == 0)
2097 if (shiftAmount >= inputWidth && inputType.
isUnsigned())
2098 return getIntAttr(getType(), APInt(0, 0,
false));
2104 value = cst->ashr(std::min(shiftAmount, inputWidth - 1));
2106 value = cst->lshr(std::min(shiftAmount, inputWidth));
2107 auto resultWidth = std::max(inputWidth - shiftAmount, 1);
2108 return getIntAttr(getType(), value.trunc(resultWidth));
2113LogicalResult ShrPrimOp::canonicalize(ShrPrimOp op, PatternRewriter &rewriter) {
2114 auto inputWidth = op.getInput().getType().base().getWidthOrSentinel();
2115 if (inputWidth <= 0)
2119 unsigned shiftAmount = op.getAmount();
2120 if (
int(shiftAmount) >= inputWidth) {
2122 if (op.getType().base().isUnsigned())
2128 shiftAmount = inputWidth - 1;
2131 replaceWithBits(op, op.getInput(), inputWidth - 1, shiftAmount, rewriter);
2135LogicalResult HeadPrimOp::canonicalize(HeadPrimOp op,
2136 PatternRewriter &rewriter) {
2137 auto inputWidth = op.getInput().getType().base().getWidthOrSentinel();
2138 if (inputWidth <= 0)
2142 unsigned keepAmount = op.getAmount();
2144 replaceWithBits(op, op.getInput(), inputWidth - 1, inputWidth - keepAmount,
2149OpFoldResult HeadPrimOp::fold(FoldAdaptor adaptor) {
2153 getInput().getType().base().getWidthOrSentinel() - getAmount();
2154 return getIntAttr(getType(), cst->lshr(shiftAmount).trunc(getAmount()));
2160OpFoldResult TailPrimOp::fold(FoldAdaptor adaptor) {
2164 cst->trunc(getType().base().getWidthOrSentinel()));
2168LogicalResult TailPrimOp::canonicalize(TailPrimOp op,
2169 PatternRewriter &rewriter) {
2170 auto inputWidth = op.getInput().getType().base().getWidthOrSentinel();
2171 if (inputWidth <= 0)
2175 unsigned dropAmount = op.getAmount();
2176 if (dropAmount !=
unsigned(inputWidth))
2182void SubaccessOp::getCanonicalizationPatterns(RewritePatternSet &results,
2184 results.add<patterns::SubaccessOfConstant>(
context);
2187OpFoldResult MultibitMuxOp::fold(FoldAdaptor adaptor) {
2189 if (adaptor.getInputs().size() == 1)
2190 return getOperand(1);
2192 if (
auto constIndex =
getConstant(adaptor.getIndex())) {
2193 auto index = constIndex->getZExtValue();
2194 if (index < getInputs().size())
2195 return getInputs()[getInputs().size() - 1 - index];
2201LogicalResult MultibitMuxOp::canonicalize(MultibitMuxOp op,
2202 PatternRewriter &rewriter) {
2206 if (llvm::all_of(op.getInputs().drop_front(), [&](
auto input) {
2207 return input == op.getInputs().front();
2215 auto indexWidth = op.getIndex().getType().getBitWidthOrSentinel();
2216 uint64_t inputSize = op.getInputs().size();
2217 if (indexWidth >= 0 && indexWidth < 64 && 1ull << indexWidth < inputSize) {
2218 rewriter.modifyOpInPlace(op, [&]() {
2219 op.getInputsMutable().erase(0, inputSize - (1ull << indexWidth));
2226 if (
auto lastSubindex = op.getInputs().back().getDefiningOp<SubindexOp>()) {
2227 if (llvm::all_of(llvm::enumerate(op.getInputs()), [&](
auto e) {
2228 auto subindex = e.value().template getDefiningOp<SubindexOp>();
2229 return subindex && lastSubindex.getInput() == subindex.getInput() &&
2230 subindex.getIndex() + e.index() + 1 == op.getInputs().size();
2232 replaceOpWithNewOpAndCopyName<SubaccessOp>(
2233 rewriter, op, lastSubindex.getInput(), op.getIndex());
2239 if (op.getInputs().size() != 2)
2243 auto uintType = op.getIndex().getType();
2244 if (uintType.getBitWidthOrSentinel() != 1)
2248 replaceOpWithNewOpAndCopyName<MuxPrimOp>(
2249 rewriter, op, op.getIndex(), op.getInputs()[0], op.getInputs()[1]);
2268 MatchingConnectOp connect;
2269 for (Operation *user : value.getUsers()) {
2271 if (isa<AttachOp, SubfieldOp, SubaccessOp, SubindexOp>(user))
2274 if (
auto aConnect = dyn_cast<FConnectLike>(user))
2275 if (aConnect.getDest() == value) {
2276 auto matchingConnect = dyn_cast<MatchingConnectOp>(*aConnect);
2279 if (!matchingConnect || (connect && connect != matchingConnect) ||
2280 matchingConnect->getBlock() != value.getParentBlock())
2282 connect = matchingConnect;
2290 PatternRewriter &rewriter) {
2293 Operation *connectedDecl = op.getDest().getDefiningOp();
2298 if (!isa<WireOp>(connectedDecl) && !isa<RegOp>(connectedDecl))
2302 cast<Forceable>(connectedDecl).isForceable())
2310 if (connectedDecl->hasOneUse())
2314 auto *declBlock = connectedDecl->getBlock();
2315 auto *srcValueOp = op.getSrc().getDefiningOp();
2318 if (!isa<WireOp>(connectedDecl))
2324 auto cnst = dyn_cast<ConstantOp>(srcValueOp);
2327 if (srcValueOp->getBlock() != declBlock)
2332 if (
auto reg = dyn_cast<RegOp>(connectedDecl))
2339 auto replacement = op.getSrc();
2342 if (srcValueOp && srcValueOp != &declBlock->front())
2343 srcValueOp->moveBefore(&declBlock->front());
2350 rewriter.eraseOp(op);
2354void ConnectOp::getCanonicalizationPatterns(RewritePatternSet &results,
2356 results.insert<patterns::ConnectExtension, patterns::ConnectSameType>(
2360LogicalResult MatchingConnectOp::canonicalize(MatchingConnectOp op,
2361 PatternRewriter &rewriter) {
2378 for (
auto *user : value.getUsers()) {
2379 auto attach = dyn_cast<AttachOp>(user);
2380 if (!attach || attach == dominatedAttach)
2382 if (attach->isBeforeInBlock(dominatedAttach))
2388LogicalResult AttachOp::canonicalize(AttachOp op, PatternRewriter &rewriter) {
2390 if (op.getNumOperands() <= 1) {
2391 rewriter.eraseOp(op);
2395 for (
auto operand : op.getOperands()) {
2402 SmallVector<Value> newOperands(op.getOperands());
2403 for (
auto newOperand : attach.getOperands())
2404 if (newOperand != operand)
2405 newOperands.push_back(newOperand);
2406 AttachOp::create(rewriter, op->getLoc(), newOperands);
2407 rewriter.eraseOp(attach);
2408 rewriter.eraseOp(op);
2416 if (
auto wire = dyn_cast_or_null<WireOp>(operand.getDefiningOp())) {
2417 if (!
hasDontTouch(wire.getOperation()) && wire->hasOneUse() &&
2418 !wire.isForceable()) {
2419 SmallVector<Value> newOperands;
2420 for (
auto newOperand : op.getOperands())
2421 if (newOperand != operand)
2422 newOperands.push_back(newOperand);
2424 AttachOp::create(rewriter, op->getLoc(), newOperands);
2425 rewriter.eraseOp(op);
2426 rewriter.eraseOp(wire);
2437 assert(llvm::hasSingleElement(region) &&
"expected single-region block");
2438 rewriter.inlineBlockBefore(®ion.front(), op, {});
2441LogicalResult WhenOp::canonicalize(WhenOp op, PatternRewriter &rewriter) {
2442 if (
auto constant = op.getCondition().getDefiningOp<firrtl::ConstantOp>()) {
2443 if (constant.getValue().isAllOnes())
2445 else if (op.hasElseRegion() && !op.getElseRegion().empty())
2448 rewriter.eraseOp(op);
2454 if (!op.getThenBlock().empty() && op.hasElseRegion() &&
2455 op.getElseBlock().empty()) {
2456 rewriter.eraseBlock(&op.getElseBlock());
2463 if (!op.getThenBlock().empty())
2467 if (!op.hasElseRegion() || op.getElseBlock().empty()) {
2468 rewriter.eraseOp(op);
2478 using OpRewritePattern::OpRewritePattern;
2479 LogicalResult matchAndRewrite(NodeOp node,
2480 PatternRewriter &rewriter)
const override {
2481 auto name = node.getNameAttr();
2482 if (!node.hasDroppableName() || node.getInnerSym() ||
2485 auto *newOp = node.getInput().getDefiningOp();
2488 rewriter.replaceOp(node, node.getInput());
2495 using OpRewritePattern::OpRewritePattern;
2496 LogicalResult matchAndRewrite(NodeOp node,
2497 PatternRewriter &rewriter)
const override {
2499 node.use_empty() || node.isForceable())
2501 rewriter.replaceAllUsesWith(node.getResult(), node.getInput());
2508template <
typename OpTy>
2510 PatternRewriter &rewriter) {
2511 if (!op.isForceable() || !op.getDataRef().use_empty())
2519LogicalResult NodeOp::fold(FoldAdaptor adaptor,
2520 SmallVectorImpl<OpFoldResult> &results) {
2529 if (!adaptor.getInput())
2532 results.push_back(adaptor.getInput());
2536void NodeOp::getCanonicalizationPatterns(RewritePatternSet &results,
2538 results.insert<FoldNodeName>(
context);
2539 results.add(demoteForceableIfUnused<NodeOp>);
2545struct AggOneShot :
public mlir::RewritePattern {
2546 AggOneShot(StringRef name, uint32_t weight, MLIRContext *
context)
2547 : RewritePattern(name, 0,
context) {}
2549 SmallVector<Value> getCompleteWrite(Operation *lhs)
const {
2550 auto lhsTy = lhs->getResult(0).getType();
2551 if (!type_isa<BundleType, FVectorType>(lhsTy))
2554 DenseMap<uint32_t, Value> fields;
2555 for (Operation *user : lhs->getResult(0).getUsers()) {
2556 if (user->getParentOp() != lhs->getParentOp())
2558 if (
auto aConnect = dyn_cast<MatchingConnectOp>(user)) {
2559 if (aConnect.getDest() == lhs->getResult(0))
2561 }
else if (
auto subField = dyn_cast<SubfieldOp>(user)) {
2562 for (Operation *subuser : subField.getResult().getUsers()) {
2563 if (
auto aConnect = dyn_cast<MatchingConnectOp>(subuser)) {
2564 if (aConnect.getDest() == subField) {
2565 if (subuser->getParentOp() != lhs->getParentOp())
2567 if (fields.count(subField.getFieldIndex()))
2569 fields[subField.getFieldIndex()] = aConnect.getSrc();
2575 }
else if (
auto subIndex = dyn_cast<SubindexOp>(user)) {
2576 for (Operation *subuser : subIndex.getResult().getUsers()) {
2577 if (
auto aConnect = dyn_cast<MatchingConnectOp>(subuser)) {
2578 if (aConnect.getDest() == subIndex) {
2579 if (subuser->getParentOp() != lhs->getParentOp())
2581 if (fields.count(subIndex.getIndex()))
2583 fields[subIndex.getIndex()] = aConnect.getSrc();
2594 SmallVector<Value> values;
2595 uint32_t total = type_isa<BundleType>(lhsTy)
2596 ? type_cast<BundleType>(lhsTy).getNumElements()
2597 : type_cast<FVectorType>(lhsTy).getNumElements();
2598 for (uint32_t i = 0; i < total; ++i) {
2599 if (!fields.count(i))
2601 values.push_back(fields[i]);
2606 LogicalResult matchAndRewrite(Operation *op,
2607 PatternRewriter &rewriter)
const override {
2608 auto values = getCompleteWrite(op);
2611 rewriter.setInsertionPointToEnd(op->getBlock());
2612 auto dest = op->getResult(0);
2613 auto destType = dest.getType();
2616 if (!type_cast<FIRRTLBaseType>(destType).isPassive())
2619 Value newVal = type_isa<BundleType>(destType)
2620 ? rewriter.createOrFold<BundleCreateOp>(op->getLoc(),
2622 : rewriter.createOrFold<VectorCreateOp>(
2623 op->
getLoc(), destType, values);
2624 rewriter.createOrFold<MatchingConnectOp>(op->getLoc(), dest, newVal);
2625 for (Operation *user : dest.getUsers()) {
2626 if (
auto subIndex = dyn_cast<SubindexOp>(user)) {
2627 for (Operation *subuser :
2628 llvm::make_early_inc_range(subIndex.getResult().getUsers()))
2629 if (auto aConnect = dyn_cast<MatchingConnectOp>(subuser))
2630 if (aConnect.getDest() == subIndex)
2631 rewriter.eraseOp(aConnect);
2632 }
else if (
auto subField = dyn_cast<SubfieldOp>(user)) {
2633 for (Operation *subuser :
2634 llvm::make_early_inc_range(subField.getResult().getUsers()))
2635 if (auto aConnect = dyn_cast<MatchingConnectOp>(subuser))
2636 if (aConnect.getDest() == subField)
2637 rewriter.eraseOp(aConnect);
2644struct WireAggOneShot :
public AggOneShot {
2645 WireAggOneShot(MLIRContext *
context)
2646 : AggOneShot(WireOp::getOperationName(), 0,
context) {}
2648struct SubindexAggOneShot :
public AggOneShot {
2649 SubindexAggOneShot(MLIRContext *
context)
2650 : AggOneShot(SubindexOp::getOperationName(), 0,
context) {}
2652struct SubfieldAggOneShot :
public AggOneShot {
2653 SubfieldAggOneShot(MLIRContext *
context)
2654 : AggOneShot(SubfieldOp::getOperationName(), 0,
context) {}
2658void WireOp::getCanonicalizationPatterns(RewritePatternSet &results,
2660 results.insert<WireAggOneShot>(
context);
2661 results.add(demoteForceableIfUnused<WireOp>);
2664void SubindexOp::getCanonicalizationPatterns(RewritePatternSet &results,
2666 results.insert<SubindexAggOneShot>(
context);
2669OpFoldResult SubindexOp::fold(FoldAdaptor adaptor) {
2670 auto attr = dyn_cast_or_null<ArrayAttr>(adaptor.getInput());
2673 return attr[getIndex()];
2676OpFoldResult SubfieldOp::fold(FoldAdaptor adaptor) {
2677 auto attr = dyn_cast_or_null<ArrayAttr>(adaptor.getInput());
2680 auto index = getFieldIndex();
2684void SubfieldOp::getCanonicalizationPatterns(RewritePatternSet &results,
2686 results.insert<SubfieldAggOneShot>(
context);
2690 ArrayRef<Attribute> operands) {
2691 for (
auto operand : operands)
2694 return ArrayAttr::get(
context, operands);
2697OpFoldResult BundleCreateOp::fold(FoldAdaptor adaptor) {
2700 if (getNumOperands() > 0)
2701 if (SubfieldOp first = getOperand(0).getDefiningOp<SubfieldOp>())
2702 if (first.getFieldIndex() == 0 &&
2703 first.getInput().getType() == getType() &&
2705 llvm::drop_begin(llvm::enumerate(getOperands())), [&](
auto elem) {
2707 elem.value().
template getDefiningOp<SubfieldOp>();
2708 return subindex && subindex.getInput() == first.getInput() &&
2709 subindex.getFieldIndex() == elem.index();
2711 return first.getInput();
2716OpFoldResult VectorCreateOp::fold(FoldAdaptor adaptor) {
2719 if (getNumOperands() > 0)
2720 if (SubindexOp first = getOperand(0).getDefiningOp<SubindexOp>())
2721 if (first.getIndex() == 0 && first.getInput().getType() == getType() &&
2723 llvm::drop_begin(llvm::enumerate(getOperands())), [&](
auto elem) {
2725 elem.value().
template getDefiningOp<SubindexOp>();
2726 return subindex && subindex.getInput() == first.getInput() &&
2727 subindex.getIndex() == elem.index();
2729 return first.getInput();
2734OpFoldResult UninferredResetCastOp::fold(FoldAdaptor adaptor) {
2735 if (getOperand().getType() == getType())
2736 return getOperand();
2744 using OpRewritePattern::OpRewritePattern;
2745 LogicalResult matchAndRewrite(RegResetOp reg,
2746 PatternRewriter &rewriter)
const override {
2748 dyn_cast_or_null<ConstantOp>(
reg.getResetValue().getDefiningOp());
2757 auto mux = dyn_cast_or_null<MuxPrimOp>(con.getSrc().getDefiningOp());
2760 auto *high = mux.getHigh().getDefiningOp();
2761 auto *low = mux.getLow().getDefiningOp();
2762 auto constOp = dyn_cast_or_null<ConstantOp>(high);
2764 if (constOp && low != reg)
2766 if (dyn_cast_or_null<ConstantOp>(low) && high == reg)
2767 constOp = dyn_cast<ConstantOp>(low);
2769 if (!constOp || constOp.getType() != reset.getType() ||
2770 constOp.getValue() != reset.getValue())
2774 auto regTy =
reg.getResult().getType();
2775 if (con.getDest().getType() != regTy || con.getSrc().getType() != regTy ||
2776 mux.getHigh().getType() != regTy || mux.getLow().getType() != regTy ||
2777 regTy.getBitWidthOrSentinel() < 0)
2788 if (constOp != &con->getBlock()->front())
2789 constOp->moveBefore(&con->getBlock()->front());
2794 rewriter.eraseOp(con);
2801 if (
auto c = v.getDefiningOp<ConstantOp>())
2802 return c.getValue().isOne();
2803 if (
auto sc = v.getDefiningOp<SpecialConstantOp>())
2804 return sc.getValue();
2813 auto resetValue = reg.getResetValue();
2814 if (reg.getType(0) != resetValue.getType())
2819 if (
auto constOp = dyn_cast_or_null<ConstantOp>(resetValue.getDefiningOp())) {
2827 (void)
dropWrite(rewriter, reg->getResult(0), {});
2828 replaceOpWithNewOpAndCopyName<NodeOp>(
2829 rewriter, reg, resetValue, reg.getNameAttr(), reg.getNameKind(),
2830 reg.getAnnotationsAttr(), reg.getInnerSymAttr(), reg.getForceable());
2834void RegResetOp::getCanonicalizationPatterns(RewritePatternSet &results,
2836 results.add<patterns::RegResetWithZeroReset, FoldResetMux>(
context);
2838 results.add(demoteForceableIfUnused<RegResetOp>);
2843 auto portTy = type_cast<BundleType>(port.getType());
2844 auto fieldIndex = portTy.getElementIndex(name);
2845 assert(fieldIndex &&
"missing field on memory port");
2848 for (
auto *op : port.getUsers()) {
2849 auto portAccess = cast<SubfieldOp>(op);
2850 if (fieldIndex != portAccess.getFieldIndex())
2855 value = conn.getSrc();
2865 auto portConst = value.getDefiningOp<ConstantOp>();
2868 return portConst.getValue().isZero();
2873 auto portTy = type_cast<BundleType>(port.getType());
2874 auto fieldIndex = portTy.getElementIndex(
data);
2875 assert(fieldIndex &&
"missing enable flag on memory port");
2877 for (
auto *op : port.getUsers()) {
2878 auto portAccess = cast<SubfieldOp>(op);
2879 if (fieldIndex != portAccess.getFieldIndex())
2881 if (!portAccess.use_empty())
2890 StringRef name, Value value) {
2891 auto portTy = type_cast<BundleType>(port.getType());
2892 auto fieldIndex = portTy.getElementIndex(name);
2893 assert(fieldIndex &&
"missing field on memory port");
2895 for (
auto *op : llvm::make_early_inc_range(port.getUsers())) {
2896 auto portAccess = cast<SubfieldOp>(op);
2897 if (fieldIndex != portAccess.getFieldIndex())
2899 rewriter.replaceAllUsesWith(portAccess, value);
2900 rewriter.eraseOp(portAccess);
2905static void erasePort(PatternRewriter &rewriter, Value port) {
2908 auto getClock = [&] {
2910 clock = SpecialConstantOp::create(rewriter, port.getLoc(),
2911 ClockType::get(rewriter.getContext()),
2920 for (
auto *op : port.getUsers()) {
2921 auto subfield = dyn_cast<SubfieldOp>(op);
2923 auto ty = port.getType();
2924 auto reg = RegOp::create(rewriter, port.getLoc(), ty, getClock());
2925 rewriter.replaceAllUsesWith(port, reg.getResult());
2934 for (
auto *accessOp : llvm::make_early_inc_range(port.getUsers())) {
2935 auto access = cast<SubfieldOp>(accessOp);
2936 for (
auto *user : llvm::make_early_inc_range(access->getUsers())) {
2937 auto connect = dyn_cast<FConnectLike>(user);
2938 if (connect && connect.getDest() == access) {
2939 rewriter.eraseOp(user);
2943 if (access.use_empty()) {
2944 rewriter.eraseOp(access);
2950 auto ty = access.getType();
2951 auto reg = RegOp::create(rewriter, access.getLoc(), ty, getClock());
2952 rewriter.replaceOp(access, reg.getResult());
2954 assert(port.use_empty() &&
"port should have no remaining uses");
2960 using OpRewritePattern::OpRewritePattern;
2961 LogicalResult matchAndRewrite(MemOp mem,
2962 PatternRewriter &rewriter)
const override {
2966 if (!firrtl::type_isa<IntType>(mem.getDataType()) ||
2967 mem.getDataType().getBitWidthOrSentinel() != 0)
2971 for (
auto port : mem.getResults())
2972 for (auto *user : port.getUsers())
2973 if (!isa<SubfieldOp>(user))
2978 for (
auto port : mem.getResults()) {
2979 for (
auto *user :
llvm::make_early_inc_range(port.getUsers())) {
2980 SubfieldOp sfop = cast<SubfieldOp>(user);
2981 StringRef fieldName = sfop.getFieldName();
2982 auto wire = replaceOpWithNewOpAndCopyName<WireOp>(
2983 rewriter, sfop, sfop.getResult().getType())
2985 if (fieldName.ends_with(
"data")) {
2987 auto zero = firrtl::ConstantOp::create(
2988 rewriter, wire.getLoc(),
2989 firrtl::type_cast<IntType>(wire.getType()), APInt::getZero(0));
2990 MatchingConnectOp::create(rewriter, wire.getLoc(), wire, zero);
2994 rewriter.eraseOp(mem);
3001 using OpRewritePattern::OpRewritePattern;
3002 LogicalResult matchAndRewrite(MemOp mem,
3003 PatternRewriter &rewriter)
const override {
3006 bool isRead =
false, isWritten =
false;
3007 for (
unsigned i = 0; i < mem.getNumResults(); ++i) {
3008 switch (mem.getPortKind(i)) {
3009 case MemOp::PortKind::Read:
3014 case MemOp::PortKind::Write:
3019 case MemOp::PortKind::Debug:
3020 case MemOp::PortKind::ReadWrite:
3023 llvm_unreachable(
"unknown port kind");
3025 assert((!isWritten || !isRead) &&
"memory is in use");
3030 if (isRead && mem.getInit())
3033 for (
auto port : mem.getResults())
3036 rewriter.eraseOp(mem);
3043 using OpRewritePattern::OpRewritePattern;
3044 LogicalResult matchAndRewrite(MemOp mem,
3045 PatternRewriter &rewriter)
const override {
3049 llvm::SmallBitVector deadPorts(mem.getNumResults());
3050 for (
auto [i, port] :
llvm::enumerate(mem.getResults())) {
3052 if (!mem.getPortAnnotation(i).empty())
3056 auto kind = mem.getPortKind(i);
3057 if (kind == MemOp::PortKind::Debug)
3066 if (kind == MemOp::PortKind::Read &&
isPortUnused(port,
"data")) {
3071 if (deadPorts.none())
3075 SmallVector<Type> resultTypes;
3076 SmallVector<StringRef> portNames;
3077 SmallVector<Attribute> portAnnotations;
3078 for (
auto [i, port] :
llvm::enumerate(mem.getResults())) {
3081 resultTypes.push_back(port.getType());
3082 portNames.push_back(mem.getPortName(i));
3083 portAnnotations.push_back(mem.getPortAnnotation(i));
3087 if (!resultTypes.empty())
3088 newOp = MemOp::create(
3089 rewriter, mem.getLoc(), resultTypes, mem.getReadLatency(),
3090 mem.getWriteLatency(), mem.getDepth(), mem.getRuw(),
3091 rewriter.getStrArrayAttr(portNames), mem.getName(), mem.getNameKind(),
3092 mem.getAnnotations(), rewriter.getArrayAttr(portAnnotations),
3093 mem.getInnerSymAttr(), mem.getInitAttr(), mem.getPrefixAttr());
3096 unsigned nextPort = 0;
3097 for (
auto [i, port] :
llvm::enumerate(mem.getResults())) {
3101 rewriter.replaceAllUsesWith(port, newOp.getResult(nextPort++));
3104 rewriter.eraseOp(mem);
3111 using OpRewritePattern::OpRewritePattern;
3112 LogicalResult matchAndRewrite(MemOp mem,
3113 PatternRewriter &rewriter)
const override {
3118 llvm::SmallBitVector deadReads(mem.getNumResults());
3119 for (
auto [i, port] :
llvm::enumerate(mem.getResults())) {
3120 if (mem.getPortKind(i) != MemOp::PortKind::ReadWrite)
3122 if (!mem.getPortAnnotation(i).empty())
3129 if (deadReads.none())
3132 SmallVector<Type> resultTypes;
3133 SmallVector<StringRef> portNames;
3134 SmallVector<Attribute> portAnnotations;
3135 for (
auto [i, port] :
llvm::enumerate(mem.getResults())) {
3137 resultTypes.push_back(
3138 MemOp::getTypeForPort(mem.getDepth(), mem.getDataType(),
3139 MemOp::PortKind::Write, mem.getMaskBits()));
3141 resultTypes.push_back(port.getType());
3143 portNames.push_back(mem.getPortName(i));
3144 portAnnotations.push_back(mem.getPortAnnotation(i));
3147 auto newOp = MemOp::create(
3148 rewriter, mem.getLoc(), resultTypes, mem.getReadLatency(),
3149 mem.getWriteLatency(), mem.getDepth(), mem.getRuw(),
3150 rewriter.getStrArrayAttr(portNames), mem.getName(), mem.getNameKind(),
3151 mem.getAnnotations(), rewriter.getArrayAttr(portAnnotations),
3152 mem.getInnerSymAttr(), mem.getInitAttr(), mem.getPrefixAttr());
3154 for (
unsigned i = 0, n = mem.getNumResults(); i < n; ++i) {
3155 auto result = mem.getResult(i);
3156 auto newResult = newOp.getResult(i);
3158 auto resultPortTy = type_cast<BundleType>(result.getType());
3162 auto replace = [&](StringRef toName, StringRef fromName) {
3163 auto fromFieldIndex = resultPortTy.getElementIndex(fromName);
3164 assert(fromFieldIndex &&
"missing enable flag on memory port");
3166 auto toField = SubfieldOp::create(rewriter, newResult.getLoc(),
3168 for (
auto *op :
llvm::make_early_inc_range(result.getUsers())) {
3169 auto fromField = cast<SubfieldOp>(op);
3170 if (fromFieldIndex != fromField.getFieldIndex())
3172 rewriter.replaceOp(fromField, toField.getResult());
3176 replace(
"addr",
"addr");
3177 replace(
"en",
"en");
3178 replace(
"clk",
"clk");
3179 replace(
"data",
"wdata");
3180 replace(
"mask",
"wmask");
3183 auto wmodeFieldIndex = resultPortTy.getElementIndex(
"wmode");
3184 for (
auto *op :
llvm::make_early_inc_range(result.getUsers())) {
3185 auto wmodeField = cast<SubfieldOp>(op);
3186 if (wmodeFieldIndex != wmodeField.getFieldIndex())
3188 rewriter.replaceOpWithNewOp<WireOp>(wmodeField, wmodeField.getType());
3191 rewriter.replaceAllUsesWith(result, newResult);
3194 rewriter.eraseOp(mem);
3201 using OpRewritePattern::OpRewritePattern;
3203 LogicalResult matchAndRewrite(MemOp mem,
3204 PatternRewriter &rewriter)
const override {
3209 const auto &summary = mem.getSummary();
3210 if (summary.isMasked || summary.isSeqMem())
3213 auto type = type_dyn_cast<IntType>(mem.getDataType());
3216 auto width = type.getBitWidthOrSentinel();
3220 llvm::SmallBitVector usedBits(width);
3221 DenseMap<unsigned, unsigned> mapping;
3226 SmallVector<BitsPrimOp> readOps;
3227 auto findReadUsers = [&](Value port, StringRef field) -> LogicalResult {
3228 auto portTy = type_cast<BundleType>(port.getType());
3229 auto fieldIndex = portTy.getElementIndex(field);
3230 assert(fieldIndex &&
"missing data port");
3232 for (
auto *op : port.getUsers()) {
3233 auto portAccess = cast<SubfieldOp>(op);
3234 if (fieldIndex != portAccess.getFieldIndex())
3237 for (
auto *user : op->getUsers()) {
3238 auto bits = dyn_cast<BitsPrimOp>(user);
3242 usedBits.set(bits.getLo(), bits.getHi() + 1);
3246 mapping[bits.getLo()] = 0;
3247 readOps.push_back(bits);
3257 SmallVector<MatchingConnectOp> writeOps;
3258 auto findWriteUsers = [&](Value port, StringRef field) -> LogicalResult {
3259 auto portTy = type_cast<BundleType>(port.getType());
3260 auto fieldIndex = portTy.getElementIndex(field);
3261 assert(fieldIndex &&
"missing data port");
3263 for (
auto *op : port.getUsers()) {
3264 auto portAccess = cast<SubfieldOp>(op);
3265 if (fieldIndex != portAccess.getFieldIndex())
3272 writeOps.push_back(conn);
3278 for (
auto [i, port] :
llvm::enumerate(mem.getResults())) {
3280 if (!mem.getPortAnnotation(i).empty())
3283 switch (mem.getPortKind(i)) {
3284 case MemOp::PortKind::Debug:
3287 case MemOp::PortKind::Write:
3288 if (failed(findWriteUsers(port,
"data")))
3291 case MemOp::PortKind::Read:
3292 if (failed(findReadUsers(port,
"data")))
3295 case MemOp::PortKind::ReadWrite:
3296 if (failed(findWriteUsers(port,
"wdata")))
3298 if (failed(findReadUsers(port,
"rdata")))
3302 llvm_unreachable(
"unknown port kind");
3306 if (usedBits.none())
3310 SmallVector<std::pair<unsigned, unsigned>> ranges;
3311 unsigned newWidth = 0;
3312 for (
int i = usedBits.find_first(); 0 <= i && i < width;) {
3313 int e = usedBits.find_next_unset(i);
3316 for (
int idx = i; idx < e; ++idx, ++newWidth) {
3317 if (
auto it = mapping.find(idx); it != mapping.end()) {
3318 it->second = newWidth;
3321 ranges.emplace_back(i, e - 1);
3322 i = e != width ? usedBits.find_next(e) : e;
3326 auto newType =
IntType::get(mem->getContext(), type.isSigned(), newWidth);
3327 SmallVector<Type> portTypes;
3328 for (
auto [i, port] :
llvm::enumerate(mem.getResults())) {
3329 portTypes.push_back(
3330 MemOp::getTypeForPort(mem.getDepth(), newType, mem.getPortKind(i)));
3332 auto newMem = rewriter.replaceOpWithNewOp<MemOp>(
3333 mem, portTypes, mem.getReadLatency(), mem.getWriteLatency(),
3334 mem.getDepth(), mem.getRuw(), mem.getPortNames(), mem.getName(),
3335 mem.getNameKind(), mem.getAnnotations(), mem.getPortAnnotations(),
3336 mem.getInnerSymAttr(), mem.getInitAttr(), mem.getPrefixAttr());
3339 auto rewriteSubfield = [&](Value port, StringRef field) {
3340 auto portTy = type_cast<BundleType>(port.getType());
3341 auto fieldIndex = portTy.getElementIndex(field);
3342 assert(fieldIndex &&
"missing data port");
3344 rewriter.setInsertionPointAfter(newMem);
3345 auto newPortAccess =
3346 SubfieldOp::create(rewriter, port.getLoc(), port, field);
3348 for (
auto *op :
llvm::make_early_inc_range(port.getUsers())) {
3349 auto portAccess = cast<SubfieldOp>(op);
3350 if (op == newPortAccess || fieldIndex != portAccess.getFieldIndex())
3352 rewriter.replaceOp(portAccess, newPortAccess.getResult());
3357 for (
auto [i, port] :
llvm::enumerate(newMem.getResults())) {
3358 switch (newMem.getPortKind(i)) {
3359 case MemOp::PortKind::Debug:
3360 llvm_unreachable(
"cannot rewrite debug port");
3361 case MemOp::PortKind::Write:
3362 rewriteSubfield(port,
"data");
3364 case MemOp::PortKind::Read:
3365 rewriteSubfield(port,
"data");
3367 case MemOp::PortKind::ReadWrite:
3368 rewriteSubfield(port,
"rdata");
3369 rewriteSubfield(port,
"wdata");
3372 llvm_unreachable(
"unknown port kind");
3376 for (
auto readOp : readOps) {
3377 rewriter.setInsertionPointAfter(readOp);
3378 auto it = mapping.find(readOp.getLo());
3379 assert(it != mapping.end() &&
"bit op mapping not found");
3382 auto newReadValue = rewriter.createOrFold<BitsPrimOp>(
3383 readOp.getLoc(), readOp.getInput(),
3384 readOp.getHi() - readOp.getLo() + it->second, it->second);
3385 rewriter.replaceAllUsesWith(readOp, newReadValue);
3386 rewriter.eraseOp(readOp);
3390 for (
auto writeOp : writeOps) {
3391 Value source = writeOp.getSrc();
3392 rewriter.setInsertionPoint(writeOp);
3394 SmallVector<Value> slices;
3395 for (
auto &[start, end] :
llvm::reverse(ranges)) {
3396 Value slice = rewriter.createOrFold<BitsPrimOp>(writeOp.getLoc(),
3397 source,
end, start);
3398 slices.push_back(slice);
3402 rewriter.createOrFold<CatPrimOp>(writeOp.getLoc(), slices);
3408 if (type.isSigned())
3410 rewriter.createOrFold<AsSIntPrimOp>(writeOp.getLoc(), catOfSlices);
3412 rewriter.replaceOpWithNewOp<MatchingConnectOp>(writeOp, writeOp.getDest(),
3422 using OpRewritePattern::OpRewritePattern;
3423 LogicalResult matchAndRewrite(MemOp mem,
3424 PatternRewriter &rewriter)
const override {
3429 auto ty = mem.getDataType();
3430 auto loc = mem.getLoc();
3431 auto *block = mem->getBlock();
3435 SmallPtrSet<Operation *, 8> connects;
3436 SmallVector<SubfieldOp> portAccesses;
3437 for (
auto [i, port] :
llvm::enumerate(mem.getResults())) {
3438 if (!mem.getPortAnnotation(i).empty())
3441 auto collect = [&, port = port](ArrayRef<StringRef> fields) {
3442 auto portTy = type_cast<BundleType>(port.getType());
3443 for (
auto field : fields) {
3444 auto fieldIndex = portTy.getElementIndex(field);
3445 assert(fieldIndex &&
"missing field on memory port");
3447 for (
auto *op : port.getUsers()) {
3448 auto portAccess = cast<SubfieldOp>(op);
3449 if (fieldIndex != portAccess.getFieldIndex())
3451 portAccesses.push_back(portAccess);
3452 for (
auto *user : portAccess->getUsers()) {
3453 auto conn = dyn_cast<FConnectLike>(user);
3456 connects.insert(conn);
3463 switch (mem.getPortKind(i)) {
3464 case MemOp::PortKind::Debug:
3466 case MemOp::PortKind::Read:
3467 if (failed(collect({
"clk",
"en",
"addr"})))
3470 case MemOp::PortKind::Write:
3471 if (failed(collect({
"clk",
"en",
"addr",
"data",
"mask"})))
3474 case MemOp::PortKind::ReadWrite:
3475 if (failed(collect({
"clk",
"en",
"addr",
"wmode",
"wdata",
"wmask"})))
3481 if (!portClock || (clock && portClock != clock))
3487 rewriter.setInsertionPointAfter(mem);
3488 auto memWire = WireOp::create(rewriter, loc, ty).getResult();
3494 rewriter.setInsertionPointToEnd(block);
3496 RegOp::create(rewriter, loc, ty, clock, mem.getName()).getResult();
3499 MatchingConnectOp::create(rewriter, loc, memWire, memReg);
3503 auto pipeline = [&](Value value, Value clock,
const Twine &name,
3505 for (
unsigned i = 0; i < latency; ++i) {
3506 std::string regName;
3508 llvm::raw_string_ostream os(regName);
3509 os << mem.getName() <<
"_" << name <<
"_" << i;
3511 auto reg = RegOp::create(rewriter, mem.getLoc(), value.getType(), clock,
3512 rewriter.getStringAttr(regName))
3514 MatchingConnectOp::create(rewriter, value.getLoc(), reg, value);
3520 const unsigned writeStages =
info.writeLatency - 1;
3525 SmallVector<std::tuple<Value, Value, Value>> writes;
3526 for (
auto [i, port] :
llvm::enumerate(mem.getResults())) {
3528 StringRef name = mem.getPortName(i);
3530 auto portPipeline = [&, port = port](StringRef field,
unsigned stages) {
3533 return pipeline(value, portClock, name +
"_" + field, stages);
3536 switch (mem.getPortKind(i)) {
3537 case MemOp::PortKind::Debug:
3538 llvm_unreachable(
"unknown port kind");
3539 case MemOp::PortKind::Read: {
3547 case MemOp::PortKind::Write: {
3548 auto data = portPipeline(
"data", writeStages);
3549 auto en = portPipeline(
"en", writeStages);
3550 auto mask = portPipeline(
"mask", writeStages);
3554 case MemOp::PortKind::ReadWrite: {
3559 auto wdata = portPipeline(
"wdata", writeStages);
3560 auto wmask = portPipeline(
"wmask", writeStages);
3565 auto wen = AndPrimOp::create(rewriter, port.getLoc(),
en,
wmode);
3567 pipeline(wen, portClock, name +
"_wen", writeStages);
3568 writes.emplace_back(
wdata, wenPipelined,
wmask);
3575 Value next = memReg;
3581 Location loc = mem.getLoc();
3582 unsigned maskGran =
info.dataWidth /
info.maskBits;
3583 SmallVector<Value> chunks;
3584 for (
unsigned i = 0; i <
info.maskBits; ++i) {
3585 unsigned hi = (i + 1) * maskGran - 1;
3586 unsigned lo = i * maskGran;
3588 auto dataPart = rewriter.createOrFold<BitsPrimOp>(loc,
data, hi, lo);
3589 auto nextPart = rewriter.createOrFold<BitsPrimOp>(loc, next, hi, lo);
3590 auto bit = rewriter.createOrFold<BitsPrimOp>(loc,
mask, i, i);
3591 auto chunk = MuxPrimOp::create(rewriter, loc, bit, dataPart, nextPart);
3592 chunks.push_back(chunk);
3595 std::reverse(chunks.begin(), chunks.end());
3596 masked = rewriter.createOrFold<CatPrimOp>(loc, chunks);
3597 next = MuxPrimOp::create(rewriter, next.getLoc(),
en, masked, next);
3599 Value typedNext = rewriter.createOrFold<BitCastOp>(next.getLoc(), ty, next);
3600 MatchingConnectOp::create(rewriter, memReg.getLoc(), memReg, typedNext);
3603 for (Operation *conn : connects)
3604 rewriter.eraseOp(
conn);
3605 for (
auto portAccess : portAccesses)
3606 rewriter.eraseOp(portAccess);
3607 rewriter.eraseOp(mem);
3614void MemOp::getCanonicalizationPatterns(RewritePatternSet &results,
3617 .insert<FoldZeroWidthMemory, FoldReadOrWriteOnlyMemory,
3618 FoldReadWritePorts, FoldUnusedPorts, FoldUnusedBits, FoldRegMems>(
3638 auto mux = dyn_cast_or_null<MuxPrimOp>(con.getSrc().getDefiningOp());
3641 auto *high = mux.getHigh().getDefiningOp();
3642 auto *low = mux.getLow().getDefiningOp();
3644 auto constOp = dyn_cast_or_null<ConstantOp>(high);
3651 bool constReg =
false;
3653 if (constOp && low == reg)
3655 else if (dyn_cast_or_null<ConstantOp>(low) && high == reg) {
3657 constOp = dyn_cast<ConstantOp>(low);
3664 if (!isa<BlockArgument>(mux.getSel()) && !constReg)
3668 auto regTy = reg.getResult().getType();
3669 if (con.getDest().getType() != regTy || con.getSrc().getType() != regTy ||
3670 mux.getHigh().getType() != regTy || mux.getLow().getType() != regTy ||
3671 regTy.getBitWidthOrSentinel() < 0)
3678 if (constReg && !
preservesInitial(reg.getInitialAttr(), constOp.getValue()))
3682 if (constOp != &con->getBlock()->front())
3683 constOp->moveBefore(&con->getBlock()->front());
3686 SmallVector<NamedAttribute, 2> attrs(reg->getDialectAttrs());
3687 auto newReg = replaceOpWithNewOpAndCopyName<RegResetOp>(
3688 rewriter, reg, reg.getResult().getType(), reg.getClockVal(),
3689 mux.getSel(), mux.getHigh(), reg.getNameAttr(), reg.getNameKindAttr(),
3690 reg.getAnnotationsAttr(), reg.getInnerSymAttr(), reg.getForceableAttr(),
3691 reg.getInitialAttr());
3692 newReg->setDialectAttrs(attrs);
3694 auto pt = rewriter.saveInsertionPoint();
3695 rewriter.setInsertionPoint(con);
3696 auto v = constReg ? (Value)constOp.getResult() : (Value)mux.getLow();
3697 replaceOpWithNewOpAndCopyName<ConnectOp>(rewriter, con, con.getDest(), v);
3698 rewriter.restoreInsertionPoint(pt);
3702LogicalResult RegOp::canonicalize(RegOp op, PatternRewriter &rewriter) {
3703 if (!
hasDontTouch(op.getOperation()) && !op.isForceable() &&
3719 PatternRewriter &rewriter,
3722 if (
auto constant = enable.getDefiningOp<firrtl::ConstantOp>()) {
3723 if (constant.getValue().isZero()) {
3724 rewriter.eraseOp(op);
3730 if (
auto constant = predicate.getDefiningOp<firrtl::ConstantOp>()) {
3731 if (constant.getValue().isZero() == eraseIfZero) {
3732 rewriter.eraseOp(op);
3740template <
class Op,
bool EraseIfZero = false>
3742 PatternRewriter &rewriter) {
3747void AssertOp::getCanonicalizationPatterns(RewritePatternSet &results,
3749 results.add(canonicalizeImmediateVerifOp<AssertOp>);
3750 results.add<patterns::AssertXWhenX>(
context);
3753void AssumeOp::getCanonicalizationPatterns(RewritePatternSet &results,
3755 results.add(canonicalizeImmediateVerifOp<AssumeOp>);
3756 results.add<patterns::AssumeXWhenX>(
context);
3759void UnclockedAssumeIntrinsicOp::getCanonicalizationPatterns(
3760 RewritePatternSet &results, MLIRContext *
context) {
3761 results.add(canonicalizeImmediateVerifOp<UnclockedAssumeIntrinsicOp>);
3762 results.add<patterns::UnclockedAssumeIntrinsicXWhenX>(
context);
3765void CoverOp::getCanonicalizationPatterns(RewritePatternSet &results,
3767 results.add(canonicalizeImmediateVerifOp<CoverOp, /* EraseIfZero = */ true>);
3774LogicalResult InvalidValueOp::canonicalize(InvalidValueOp op,
3775 PatternRewriter &rewriter) {
3777 if (op.use_empty()) {
3778 rewriter.eraseOp(op);
3785 if (op->hasOneUse() &&
3786 (isa<BitsPrimOp, HeadPrimOp, ShrPrimOp, TailPrimOp, SubfieldOp,
3787 SubindexOp, AsSIntPrimOp, AsUIntPrimOp, NotPrimOp, BitCastOp>(
3788 *op->user_begin()) ||
3789 (isa<CvtPrimOp>(*op->user_begin()) &&
3790 type_isa<SIntType>(op->user_begin()->getOperand(0).getType())) ||
3791 (isa<AndRPrimOp, XorRPrimOp, OrRPrimOp>(*op->user_begin()) &&
3792 type_cast<FIRRTLBaseType>(op->user_begin()->getOperand(0).getType())
3793 .getBitWidthOrSentinel() > 0))) {
3794 auto *modop = *op->user_begin();
3795 auto inv = InvalidValueOp::create(rewriter, op.getLoc(),
3796 modop->getResult(0).getType());
3797 rewriter.replaceAllOpUsesWith(modop, inv);
3798 rewriter.eraseOp(modop);
3799 rewriter.eraseOp(op);
3805OpFoldResult InvalidValueOp::fold(FoldAdaptor adaptor) {
3806 if (getType().getBitWidthOrSentinel() == 0 && isa<IntType>(getType()))
3807 return getIntAttr(getType(), APInt(0, 0, isa<SIntType>(getType())));
3815OpFoldResult ClockGateIntrinsicOp::fold(FoldAdaptor adaptor) {
3824 return BoolAttr::get(getContext(),
false);
3828 return BoolAttr::get(getContext(),
false);
3833LogicalResult ClockGateIntrinsicOp::canonicalize(ClockGateIntrinsicOp op,
3834 PatternRewriter &rewriter) {
3836 if (
auto testEnable = op.getTestEnable()) {
3837 if (
auto constOp = testEnable.getDefiningOp<ConstantOp>()) {
3838 if (constOp.getValue().isZero()) {
3839 rewriter.modifyOpInPlace(op,
3840 [&] { op.getTestEnableMutable().clear(); });
3856 auto forceable = op.getRef().getDefiningOp<Forceable>();
3857 if (!forceable || !forceable.isForceable() ||
3858 op.getRef() != forceable.getDataRef() ||
3859 op.getType() != forceable.getDataType())
3861 rewriter.replaceAllUsesWith(op, forceable.getData());
3865void RefResolveOp::getCanonicalizationPatterns(RewritePatternSet &results,
3867 results.insert<patterns::RefResolveOfRefSend>(
context);
3871OpFoldResult RefCastOp::fold(FoldAdaptor adaptor) {
3873 if (getInput().getType() == getType())
3879 auto constOp = operand.getDefiningOp<ConstantOp>();
3880 return constOp && constOp.getValue().isZero();
3883template <
typename Op>
3886 rewriter.eraseOp(op);
3892void RefForceOp::getCanonicalizationPatterns(RewritePatternSet &results,
3894 results.add(eraseIfPredFalse<RefForceOp>);
3896void RefForceInitialOp::getCanonicalizationPatterns(RewritePatternSet &results,
3898 results.add(eraseIfPredFalse<RefForceInitialOp>);
3900void RefReleaseOp::getCanonicalizationPatterns(RewritePatternSet &results,
3902 results.add(eraseIfPredFalse<RefReleaseOp>);
3904void RefReleaseInitialOp::getCanonicalizationPatterns(
3905 RewritePatternSet &results, MLIRContext *
context) {
3906 results.add(eraseIfPredFalse<RefReleaseInitialOp>);
3913OpFoldResult HasBeenResetIntrinsicOp::fold(FoldAdaptor adaptor) {
3919 if (adaptor.getReset())
3924 if (
isUInt1(getReset().getType()) && adaptor.getClock())
3937 [&](
auto ty) ->
bool {
return isTypeEmpty(ty.getElementType()); })
3938 .Case<BundleType>([&](
auto ty) ->
bool {
3939 for (
auto elem : ty.getElements())
3944 .Case<IntType>([&](
auto ty) {
return ty.getWidth() == 0; })
3945 .Default([](
auto) ->
bool {
return false; });
3948LogicalResult FPGAProbeIntrinsicOp::canonicalize(FPGAProbeIntrinsicOp op,
3949 PatternRewriter &rewriter) {
3950 auto firrtlTy = type_dyn_cast<FIRRTLType>(op.getInput().getType());
3957 rewriter.eraseOp(op);
3965LogicalResult LayerBlockOp::canonicalize(LayerBlockOp op,
3966 PatternRewriter &rewriter) {
3969 if (op.getBody()->empty()) {
3970 rewriter.eraseOp(op);
3981OpFoldResult UnsafeDomainCastOp::fold(FoldAdaptor adaptor) {
3983 if (getDomains().
empty())
assert(baseType &&"element must be base type")
static std::unique_ptr< Context > context
static bool hasKnownWidthIntTypes(Operation *op)
Return true if this operation's operands and results all have a known width.
static LogicalResult canonicalizeImmediateVerifOp(Op op, PatternRewriter &rewriter)
static bool isDefinedByOneConstantOp(Value v)
static Attribute collectFields(MLIRContext *context, ArrayRef< Attribute > operands)
static LogicalResult canonicalizeSingleSetConnect(MatchingConnectOp op, PatternRewriter &rewriter)
static void erasePort(PatternRewriter &rewriter, Value port)
static void replaceOpWithRegion(PatternRewriter &rewriter, Operation *op, Region ®ion)
Replaces the given op with the contents of the given single-block region.
static std::optional< APSInt > getExtendedConstant(Value operand, Attribute constant, int32_t destWidth)
Implicitly replace the operand to a constant folding operation with a const 0 in case the operand is ...
static Value getPortFieldValue(Value port, StringRef name)
static AttachOp getDominatingAttachUser(Value value, AttachOp dominatedAttach)
If the specified value has an AttachOp user strictly dominating by "dominatingAttach" then return it.
static OpTy replaceOpWithNewOpAndCopyName(PatternRewriter &rewriter, Operation *op, Args &&...args)
A wrapper of PatternRewriter::replaceOpWithNewOp to propagate "name" attribute.
static void updateName(PatternRewriter &rewriter, Operation *op, StringAttr name)
Set the name of an op based on the best of two names: The current name, and the name passed in.
static bool isTypeEmpty(FIRRTLType type)
static bool isUInt1(Type type)
Return true if this value is 1 bit UInt.
static LogicalResult demoteForceableIfUnused(OpTy op, PatternRewriter &rewriter)
static bool isPortDisabled(Value port)
static LogicalResult eraseIfZeroOrNotZero(Operation *op, Value predicate, Value enable, PatternRewriter &rewriter, bool eraseIfZero)
static APInt getMaxSignedValue(unsigned bitWidth)
Get the largest signed value of a given bit width.
static Value dropWrite(PatternRewriter &rewriter, OpResult old, Value passthrough)
static LogicalResult canonicalizePrimOp(Operation *op, PatternRewriter &rewriter, const function_ref< OpFoldResult(ArrayRef< Attribute >)> &canonicalize)
Applies the canonicalization function canonicalize to the given operation.
static void replaceWithBits(Operation *op, Value value, unsigned hiBit, unsigned loBit, PatternRewriter &rewriter)
Replace the specified operation with a 'bits' op from the specified hi/lo bits.
static std::optional< bool > getBoolValue(Attribute attr)
static LogicalResult canonicalizeRegResetWithOneReset(RegResetOp reg, PatternRewriter &rewriter)
static LogicalResult eraseIfPredFalse(Op op, PatternRewriter &rewriter)
static OpFoldResult foldMux(OpTy op, typename OpTy::FoldAdaptor adaptor)
static APInt getMaxUnsignedValue(unsigned bitWidth)
Get the largest unsigned value of a given bit width.
static std::optional< APSInt > getConstant(Attribute operand)
Determine the value of a constant operand for the sake of constant folding.
static void replacePortField(PatternRewriter &rewriter, Value port, StringRef name, Value value)
BinOpKind
This is the policy for folding, which depends on the sort of operator we're processing.
static bool isPortUnused(Value port, StringRef data)
static bool isOkToPropagateName(Operation *op)
static LogicalResult canonicalizeRefResolveOfForceable(RefResolveOp op, PatternRewriter &rewriter)
static Attribute constFoldFIRRTLBinaryOp(Operation *op, ArrayRef< Attribute > operands, BinOpKind opKind, const function_ref< APInt(const APSInt &, const APSInt &)> &calculate)
Applies the constant folding function calculate to the given operands.
static APInt getMinSignedValue(unsigned bitWidth)
Get the smallest signed value of a given bit width.
static LogicalResult foldHiddenReset(RegOp reg, PatternRewriter &rewriter)
static Value moveNameHint(OpResult old, Value passthrough)
static void replaceOpAndCopyName(PatternRewriter &rewriter, Operation *op, Value newValue)
A wrapper of PatternRewriter::replaceOp to propagate "name" attribute.
static Location getLoc(DefSlot slot)
static InstancePath empty
AndRCat(MLIRContext *context)
bool handleConstant(mlir::PatternRewriter &rewriter, Operation *op, ConstantOp value, SmallVectorImpl< Value > &remaining) const override
Handle a constant operand in the cat operation.
bool getIdentityValue() const override
Return the unit value for this reduction operation:
bool handleConstant(mlir::PatternRewriter &rewriter, Operation *op, ConstantOp value, SmallVectorImpl< Value > &remaining) const override
Handle a constant operand in the cat operation.
OrRCat(MLIRContext *context)
bool getIdentityValue() const override
Return the unit value for this reduction operation:
virtual bool getIdentityValue() const =0
Return the unit value for this reduction operation:
virtual bool handleConstant(mlir::PatternRewriter &rewriter, Operation *op, ConstantOp constantOp, SmallVectorImpl< Value > &remaining) const =0
Handle a constant operand in the cat operation.
LogicalResult matchAndRewrite(Operation *op, mlir::PatternRewriter &rewriter) const override
ReductionCat(MLIRContext *context, llvm::StringLiteral opName)
XorRCat(MLIRContext *context)
bool handleConstant(mlir::PatternRewriter &rewriter, Operation *op, ConstantOp value, SmallVectorImpl< Value > &remaining) const override
Handle a constant operand in the cat operation.
bool getIdentityValue() const override
Return the unit value for this reduction operation:
This class provides a read-only projection over the MLIR attributes that represent a set of annotatio...
This class implements the same functionality as TypeSwitch except that it uses firrtl::type_dyn_cast ...
FIRRTLTypeSwitch< T, ResultT > & Case(CallableT &&caseFn)
Add a case on the given type.
This is the common base class between SIntType and UIntType.
int32_t getWidthOrSentinel() const
Return the width of this type, or -1 if it has none specified.
static IntType get(MLIRContext *context, bool isSigned, int32_t widthOrSentinel=-1, bool isConst=false)
Return an SIntType or UIntType with the specified signedness, width, and constness.
bool hasWidth() const
Return true if this integer type has a known width.
std::optional< int32_t > getWidth() const
Return an optional containing the width, if the width is known (or empty if width is unknown).
uint64_t getWidth(Type t)
Forceable replaceWithNewForceability(Forceable op, bool forceable, ::mlir::PatternRewriter *rewriter=nullptr)
Replace a Forceable op with equivalent, changing whether forceable.
bool areAnonymousTypesEquivalent(FIRRTLBaseType lhs, FIRRTLBaseType rhs)
Return true if anonymous types of given arguments are equivalent by pointer comparison.
IntegerAttr getIntAttr(Type type, const APInt &value)
Utiility for generating a constant attribute.
bool hasDontTouch(Value value)
Check whether a block argument ("port") or the operation defining a value has a DontTouch annotation,...
bool hasDroppableName(Operation *op)
Return true if the name is droppable.
bool preservesInitial(IntegerAttr initial, std::optional< APInt > foldedValue=std::nullopt)
Return true if replacing a register carrying the time-zero initial value with foldedValue does not ch...
MatchingConnectOp getSingleConnectUserOf(Value value)
Scan all the uses of the specified value, checking to see if there is exactly one connect that has th...
std::optional< int64_t > getBitWidth(FIRRTLBaseType type, bool ignoreFlip=false)
IntegerAttr getIntZerosAttr(Type type)
Utility for generating a constant zero attribute.
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.
APSInt extOrTruncZeroWidth(APSInt value, unsigned width)
A safe version of APSInt::extOrTrunc that will NOT assert on zero-width signed APSInts.
APInt sextZeroWidth(APInt value, unsigned width)
A safe version of APInt::sext that will NOT assert on zero-width signed APSInts.
StringRef chooseName(StringRef a, StringRef b)
Choose a good name for an item from two options.
static bool isConstantZero(Attribute operand)
Determine whether a constant operand is a zero value.
static bool isConstantOne(Attribute operand)
Determine whether a constant operand is a one value.
reg(value, clock, reset=None, reset_value=None, name=None, sym_name=None)
LogicalResult matchAndRewrite(BitsPrimOp bits, mlir::PatternRewriter &rewriter) const override