CIRCT 24.0.0git
Loading...
Searching...
No Matches
FIRRTLFolds.cpp
Go to the documentation of this file.
1//===- FIRRTLFolds.cpp - Implement folds and canonicalizations for ops ----===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implement the folding and canonicalizations for FIRRTL ops.
10//
11//===----------------------------------------------------------------------===//
12
17#include "circt/Support/APInt.h"
18#include "circt/Support/LLVM.h"
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"
27
28using namespace circt;
29using namespace firrtl;
30
31// Drop writes to old and pass through passthrough to make patterns easier to
32// write.
33static Value dropWrite(PatternRewriter &rewriter, OpResult old,
34 Value passthrough) {
35 SmallPtrSet<Operation *, 8> users;
36 for (auto *user : old.getUsers())
37 users.insert(user);
38 for (Operation *user : users)
39 if (auto connect = dyn_cast<FConnectLike>(user))
40 if (connect.getDest() == old)
41 rewriter.eraseOp(user);
42 return passthrough;
43}
44
45// Return true if it is OK to propagate the name to the operation.
46// Non-pure operations such as instances, registers, and memories are not
47// allowed to update names for name stabilities and LEC reasons.
48static bool isOkToPropagateName(Operation *op) {
49 // Conservatively disallow operations with regions to prevent performance
50 // regression due to recursive calls to sub-regions in mlir::isPure.
51 if (op->getNumRegions() != 0)
52 return false;
53 return mlir::isPure(op) || isa<NodeOp, WireOp>(op);
54}
55
56// Move a name hint from a soon to be deleted operation to a new operation.
57// Pass through the new operation to make patterns easier to write. This cannot
58// move a name to a port (block argument), doing so would require rewriting all
59// instance sites as well as the module.
60static Value moveNameHint(OpResult old, Value passthrough) {
61 Operation *op = passthrough.getDefiningOp();
62 // This should handle ports, but it isn't clear we can change those in
63 // canonicalizers.
64 assert(op && "passthrough must be an operation");
65 Operation *oldOp = old.getOwner();
66 auto name = oldOp->getAttrOfType<StringAttr>("name");
67 if (name && !name.getValue().empty() && isOkToPropagateName(op))
68 op->setAttr("name", name);
69 return passthrough;
70}
71
72// Declarative canonicalization patterns
73namespace circt {
74namespace firrtl {
75namespace patterns {
76#include "circt/Dialect/FIRRTL/FIRRTLCanonicalization.h.inc"
77} // namespace patterns
78} // namespace firrtl
79} // namespace circt
80
81/// Return true if this operation's operands and results all have a known width.
82/// This only works for integer types.
83static bool hasKnownWidthIntTypes(Operation *op) {
84 auto resultType = type_cast<IntType>(op->getResult(0).getType());
85 if (!resultType.hasWidth())
86 return false;
87 for (Value operand : op->getOperands())
88 if (!type_cast<IntType>(operand.getType()).hasWidth())
89 return false;
90 return true;
91}
92
93/// Return true if this value is 1 bit UInt.
94static bool isUInt1(Type type) {
95 auto t = type_dyn_cast<UIntType>(type);
96 if (!t || !t.hasWidth() || t.getWidth() != 1)
97 return false;
98 return true;
99}
100
101/// Set the name of an op based on the best of two names: The current name, and
102/// the name passed in.
103static void updateName(PatternRewriter &rewriter, Operation *op,
104 StringAttr name) {
105 // Should never rename InstanceOp
106 if (!name || name.getValue().empty() || !isOkToPropagateName(op))
107 return;
108 assert((!isa<InstanceOp, RegOp, RegResetOp>(op)) && "Should never rename");
109 auto newName = name.getValue(); // old name is interesting
110 auto newOpName = op->getAttrOfType<StringAttr>("name");
111 // new name might not be interesting
112 if (newOpName)
113 newName = chooseName(newOpName.getValue(), name.getValue());
114 // Only update if needed
115 if (!newOpName || newOpName.getValue() != newName)
116 rewriter.modifyOpInPlace(
117 op, [&] { op->setAttr("name", rewriter.getStringAttr(newName)); });
118}
119
120/// A wrapper of `PatternRewriter::replaceOp` to propagate "name" attribute.
121/// If a replaced op has a "name" attribute, this function propagates the name
122/// to the new value.
123static void replaceOpAndCopyName(PatternRewriter &rewriter, Operation *op,
124 Value newValue) {
125 if (auto *newOp = newValue.getDefiningOp()) {
126 auto name = op->getAttrOfType<StringAttr>("name");
127 updateName(rewriter, newOp, name);
128 }
129 rewriter.replaceOp(op, newValue);
130}
131
132/// A wrapper of `PatternRewriter::replaceOpWithNewOp` to propagate "name"
133/// attribute. If a replaced op has a "name" attribute, this function propagates
134/// the name to the new value.
135template <typename OpTy, typename... Args>
136static OpTy replaceOpWithNewOpAndCopyName(PatternRewriter &rewriter,
137 Operation *op, Args &&...args) {
138 auto name = op->getAttrOfType<StringAttr>("name");
139 auto newOp =
140 rewriter.replaceOpWithNewOp<OpTy>(op, std::forward<Args>(args)...);
141 updateName(rewriter, newOp, name);
142 return newOp;
143}
144
145/// Return true if the name is droppable. Note that this is different from
146/// `isUselessName` because non-useless names may be also droppable.
148 if (auto namableOp = dyn_cast<firrtl::FNamableOp>(op))
149 return namableOp.hasDroppableName();
150 return false;
151}
152
153/// Implicitly replace the operand to a constant folding operation with a const
154/// 0 in case the operand is non-constant but has a bit width 0, or if the
155/// operand is an invalid value.
156///
157/// This makes constant folding significantly easier, as we can simply pass the
158/// operands to an operation through this function to appropriately replace any
159/// zero-width dynamic values or invalid values with a constant of value 0.
160static std::optional<APSInt>
161getExtendedConstant(Value operand, Attribute constant, int32_t destWidth) {
162 assert(type_cast<IntType>(operand.getType()) &&
163 "getExtendedConstant is limited to integer types");
164
165 // We never support constant folding to unknown width values.
166 if (destWidth < 0)
167 return {};
168
169 // Extension signedness follows the operand sign.
170 if (IntegerAttr result = dyn_cast_or_null<IntegerAttr>(constant))
171 return extOrTruncZeroWidth(result.getAPSInt(), destWidth);
172
173 // If the operand is zero bits, then we can return a zero of the result
174 // type.
175 if (type_cast<IntType>(operand.getType()).getWidth() == 0)
176 return APSInt(destWidth,
177 type_cast<IntType>(operand.getType()).isUnsigned());
178 return {};
179}
180
181/// Determine the value of a constant operand for the sake of constant folding.
182static std::optional<APSInt> getConstant(Attribute operand) {
183 if (!operand)
184 return {};
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();
189 return {};
190}
191
192/// Determine whether a constant operand is a zero value for the sake of
193/// constant folding. This considers `invalidvalue` to be zero.
194static bool isConstantZero(Attribute operand) {
195 if (auto cst = getConstant(operand))
196 return cst->isZero();
197 return false;
198}
199
200/// Determine whether a constant operand is a one value for the sake of constant
201/// folding.
202static bool isConstantOne(Attribute operand) {
203 if (auto cst = getConstant(operand))
204 return cst->isOne();
205 return false;
206}
207
208/// This is the policy for folding, which depends on the sort of operator we're
209/// processing.
210enum class BinOpKind {
211 Normal,
212 Compare,
214};
215
216/// Applies the constant folding function `calculate` to the given operands.
217///
218/// Sign or zero extends the operands appropriately to the bitwidth of the
219/// result type if \p useDstWidth is true, else to the larger of the two operand
220/// bit widths and depending on whether the operation is to be performed on
221/// signed or unsigned operands.
222static Attribute constFoldFIRRTLBinaryOp(
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");
226
227 // We cannot fold something to an unknown width.
228 auto resultType = type_cast<IntType>(op->getResult(0).getType());
229 if (resultType.getWidthOrSentinel() < 0)
230 return {};
231
232 // Any binary op returning i0 is 0.
233 if (resultType.getWidthOrSentinel() == 0)
234 return getIntAttr(resultType, APInt(0, 0, resultType.isSigned()));
235
236 // Determine the operand widths. This is either dictated by the operand type,
237 // or if that type is an unsized integer, by the actual bits necessary to
238 // represent the constant value.
239 auto lhsWidth =
240 type_cast<IntType>(op->getOperand(0).getType()).getWidthOrSentinel();
241 auto rhsWidth =
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());
247
248 // Compares extend the operands to the widest of the operand types, not to the
249 // result type.
250 int32_t operandWidth;
251 switch (opKind) {
253 operandWidth = resultType.getWidthOrSentinel();
254 break;
256 // Compares compute with the widest operand, not at the destination type
257 // (which is always i1).
258 operandWidth = std::max(1, std::max(lhsWidth, rhsWidth));
259 break;
261 operandWidth =
262 std::max(std::max(lhsWidth, rhsWidth), resultType.getWidthOrSentinel());
263 break;
264 }
265
266 auto lhs = getExtendedConstant(op->getOperand(0), operands[0], operandWidth);
267 if (!lhs)
268 return {};
269 auto rhs = getExtendedConstant(op->getOperand(1), operands[1], operandWidth);
270 if (!rhs)
271 return {};
272
273 APInt resultValue = calculate(*lhs, *rhs);
274
275 // If the result type is smaller than the computation then we need to
276 // narrow the constant after the calculation.
277 if (opKind == BinOpKind::DivideOrShift)
278 resultValue = resultValue.trunc(resultType.getWidthOrSentinel());
279
280 assert((unsigned)resultType.getWidthOrSentinel() ==
281 resultValue.getBitWidth());
282 return getIntAttr(resultType, resultValue);
283}
284
285/// Applies the canonicalization function `canonicalize` to the given operation.
286///
287/// Determines which (if any) of the operation's operands are constants, and
288/// provides them as arguments to the callback function. Any `invalidvalue` in
289/// the input is mapped to a constant zero. The value returned from the callback
290/// is used as the replacement for `op`, and an additional pad operation is
291/// inserted if necessary. Does nothing if the result of `op` is of unknown
292/// width, in which case the necessity of a pad cannot be determined.
293static LogicalResult canonicalizePrimOp(
294 Operation *op, PatternRewriter &rewriter,
295 const function_ref<OpFoldResult(ArrayRef<Attribute>)> &canonicalize) {
296 // Can only operate on FIRRTL primitive operations.
297 if (op->getNumResults() != 1)
298 return failure();
299 auto type = type_dyn_cast<FIRRTLBaseType>(op->getResult(0).getType());
300 if (!type)
301 return failure();
302
303 // Can only operate on operations with a known result width.
304 auto width = type.getBitWidthOrSentinel();
305 if (width < 0)
306 return failure();
307
308 // Determine which of the operands are constants.
309 SmallVector<Attribute, 3> constOperands;
310 constOperands.reserve(op->getNumOperands());
311 for (auto operand : op->getOperands()) {
312 Attribute attr;
313 if (auto *defOp = operand.getDefiningOp())
314 TypeSwitch<Operation *>(defOp).Case<ConstantOp, SpecialConstantOp>(
315 [&](auto op) { attr = op.getValueAttr(); });
316 constOperands.push_back(attr);
317 }
318
319 // Perform the canonicalization and materialize the result if it is a
320 // constant.
321 auto result = canonicalize(constOperands);
322 if (!result)
323 return failure();
324 Value resultValue;
325 if (auto cst = dyn_cast<Attribute>(result))
326 resultValue = op->getDialect()
327 ->materializeConstant(rewriter, cst, type, op->getLoc())
328 ->getResult(0);
329 else
330 resultValue = cast<Value>(result);
331
332 // Insert a pad if the type widths disagree.
333 if (width !=
334 type_cast<FIRRTLBaseType>(resultValue.getType()).getBitWidthOrSentinel())
335 resultValue = PadPrimOp::create(rewriter, op->getLoc(), resultValue, width);
336
337 // Insert a cast if this is a uint vs. sint or vice versa.
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);
343
344 assert(type == resultValue.getType() && "canonicalization changed type");
345 replaceOpAndCopyName(rewriter, op, resultValue);
346 return success();
347}
348
349/// Get the largest unsigned value of a given bit width. Returns a 1-bit zero
350/// value if `bitWidth` is 0.
351static APInt getMaxUnsignedValue(unsigned bitWidth) {
352 return bitWidth > 0 ? APInt::getMaxValue(bitWidth) : APInt();
353}
354
355/// Get the smallest signed value of a given bit width. Returns a 1-bit zero
356/// value if `bitWidth` is 0.
357static APInt getMinSignedValue(unsigned bitWidth) {
358 return bitWidth > 0 ? APInt::getSignedMinValue(bitWidth) : APInt();
359}
360
361/// Get the largest signed value of a given bit width. Returns a 1-bit zero
362/// value if `bitWidth` is 0.
363static APInt getMaxSignedValue(unsigned bitWidth) {
364 return bitWidth > 0 ? APInt::getSignedMaxValue(bitWidth) : APInt();
365}
366
367//===----------------------------------------------------------------------===//
368// Fold Hooks
369//===----------------------------------------------------------------------===//
370
371OpFoldResult ConstantOp::fold(FoldAdaptor adaptor) {
372 assert(adaptor.getOperands().empty() && "constant has no operands");
373 return getValueAttr();
374}
375
376OpFoldResult SpecialConstantOp::fold(FoldAdaptor adaptor) {
377 assert(adaptor.getOperands().empty() && "constant has no operands");
378 return getValueAttr();
379}
380
381OpFoldResult AggregateConstantOp::fold(FoldAdaptor adaptor) {
382 assert(adaptor.getOperands().empty() && "constant has no operands");
383 return getFieldsAttr();
384}
385
386OpFoldResult StringConstantOp::fold(FoldAdaptor adaptor) {
387 assert(adaptor.getOperands().empty() && "constant has no operands");
388 return getValueAttr();
389}
390
391OpFoldResult FIntegerConstantOp::fold(FoldAdaptor adaptor) {
392 assert(adaptor.getOperands().empty() && "constant has no operands");
393 return getValueAttr();
394}
395
396OpFoldResult BoolConstantOp::fold(FoldAdaptor adaptor) {
397 assert(adaptor.getOperands().empty() && "constant has no operands");
398 return getValueAttr();
399}
400
401OpFoldResult DoubleConstantOp::fold(FoldAdaptor adaptor) {
402 assert(adaptor.getOperands().empty() && "constant has no operands");
403 return getValueAttr();
404}
405
406//===----------------------------------------------------------------------===//
407// Binary Operators
408//===----------------------------------------------------------------------===//
409
410OpFoldResult AddPrimOp::fold(FoldAdaptor adaptor) {
412 *this, adaptor.getOperands(), BinOpKind::Normal,
413 [=](const APSInt &a, const APSInt &b) { return a + b; });
414}
415
416void AddPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
417 MLIRContext *context) {
418 results.insert<patterns::moveConstAdd, patterns::AddOfZero,
419 patterns::AddOfSelf, patterns::AddOfPad>(context);
420}
421
422OpFoldResult SubPrimOp::fold(FoldAdaptor adaptor) {
424 *this, adaptor.getOperands(), BinOpKind::Normal,
425 [=](const APSInt &a, const APSInt &b) { return a - b; });
426}
427
428void SubPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
429 MLIRContext *context) {
430 results.insert<patterns::SubOfZero, patterns::SubFromZeroSigned,
431 patterns::SubFromZeroUnsigned, patterns::SubOfSelf,
432 patterns::SubOfPadL, patterns::SubOfPadR>(context);
433}
434
435OpFoldResult MulPrimOp::fold(FoldAdaptor adaptor) {
436 // mul(x, 0) -> 0
437 //
438 // This is legal because it aligns with the Scala FIRRTL Compiler
439 // interpretation of lowering invalid to constant zero before constant
440 // propagation. Note: the Scala FIRRTL Compiler does NOT currently optimize
441 // multiplication this way and will emit "x * 0".
442 if (isConstantZero(adaptor.getRhs()) || isConstantZero(adaptor.getLhs()))
443 return getIntZerosAttr(getType());
444
446 *this, adaptor.getOperands(), BinOpKind::Normal,
447 [=](const APSInt &a, const APSInt &b) { return a * b; });
448}
449
450OpFoldResult DivPrimOp::fold(FoldAdaptor adaptor) {
451 /// div(x, x) -> 1
452 ///
453 /// Division by zero is undefined in the FIRRTL specification. This fold
454 /// exploits that fact to optimize self division to one. Note: this should
455 /// supersede any division with invalid or zero. Division of invalid by
456 /// invalid should be one.
457 if (getLhs() == getRhs()) {
458 auto width = getType().base().getWidthOrSentinel();
459 if (width == -1)
460 width = 2;
461 // Only fold if we have at least 1 bit of width to represent the `1` value.
462 if (width != 0)
463 return getIntAttr(getType(), APInt(width, 1));
464 }
465
466 // div(0, x) -> 0
467 //
468 // This is legal because it aligns with the Scala FIRRTL Compiler
469 // interpretation of lowering invalid to constant zero before constant
470 // propagation. Note: the Scala FIRRTL Compiler does NOT currently optimize
471 // division this way and will emit "0 / x".
472 if (isConstantZero(adaptor.getLhs()) && !isConstantZero(adaptor.getRhs()))
473 return getIntZerosAttr(getType());
474
475 /// div(x, 1) -> x : (uint, uint) -> uint
476 ///
477 /// UInt division by one returns the numerator. SInt division can't
478 /// be folded here because it increases the return type bitwidth by
479 /// one and requires sign extension (a new op).
480 if (auto rhsCst = dyn_cast_or_null<IntegerAttr>(adaptor.getRhs()))
481 if (rhsCst.getValue().isOne() && getLhs().getType() == getType())
482 return getLhs();
483
485 *this, adaptor.getOperands(), BinOpKind::DivideOrShift,
486 [=](const APSInt &a, const APSInt &b) -> APInt {
487 if (!!b)
488 return a / b;
489 return APInt(a.getBitWidth(), 0);
490 });
491}
492
493OpFoldResult RemPrimOp::fold(FoldAdaptor adaptor) {
494 // rem(x, x) -> 0
495 //
496 // Division by zero is undefined in the FIRRTL specification. This fold
497 // exploits that fact to optimize self division remainder to zero. Note:
498 // this should supersede any division with invalid or zero. Remainder of
499 // division of invalid by invalid should be zero.
500 if (getLhs() == getRhs())
501 return getIntZerosAttr(getType());
502
503 // rem(0, x) -> 0
504 //
505 // This is legal because it aligns with the Scala FIRRTL Compiler
506 // interpretation of lowering invalid to constant zero before constant
507 // propagation. Note: the Scala FIRRTL Compiler does NOT currently optimize
508 // division this way and will emit "0 % x".
509 if (isConstantZero(adaptor.getLhs()))
510 return getIntZerosAttr(getType());
511
513 *this, adaptor.getOperands(), BinOpKind::DivideOrShift,
514 [=](const APSInt &a, const APSInt &b) -> APInt {
515 if (!!b)
516 return a % b;
517 return APInt(a.getBitWidth(), 0);
518 });
519}
520
521OpFoldResult DShlPrimOp::fold(FoldAdaptor adaptor) {
523 *this, adaptor.getOperands(), BinOpKind::DivideOrShift,
524 [=](const APSInt &a, const APSInt &b) -> APInt { return a.shl(b); });
525}
526
527OpFoldResult DShlwPrimOp::fold(FoldAdaptor adaptor) {
529 *this, adaptor.getOperands(), BinOpKind::DivideOrShift,
530 [=](const APSInt &a, const APSInt &b) -> APInt { return a.shl(b); });
531}
532
533OpFoldResult DShrPrimOp::fold(FoldAdaptor adaptor) {
535 *this, adaptor.getOperands(), BinOpKind::DivideOrShift,
536 [=](const APSInt &a, const APSInt &b) -> APInt {
537 return getType().base().isUnsigned() || !a.getBitWidth() ? a.lshr(b)
538 : a.ashr(b);
539 });
540}
541
542// TODO: Move to DRR.
543OpFoldResult AndPrimOp::fold(FoldAdaptor adaptor) {
544 if (auto rhsCst = getConstant(adaptor.getRhs())) {
545 /// and(x, 0) -> 0, 0 is largest or is implicit zero extended
546 if (rhsCst->isZero())
547 return getIntZerosAttr(getType());
548
549 /// and(x, -1) -> x
550 if (rhsCst->isAllOnes() && getLhs().getType() == getType() &&
551 getRhs().getType() == getType())
552 return getLhs();
553 }
554
555 if (auto lhsCst = getConstant(adaptor.getLhs())) {
556 /// and(0, x) -> 0, 0 is largest or is implicit zero extended
557 if (lhsCst->isZero())
558 return getIntZerosAttr(getType());
559
560 /// and(-1, x) -> x
561 if (lhsCst->isAllOnes() && getLhs().getType() == getType() &&
562 getRhs().getType() == getType())
563 return getRhs();
564 }
565
566 /// and(x, x) -> x
567 if (getLhs() == getRhs() && getRhs().getType() == getType())
568 return getRhs();
569
571 *this, adaptor.getOperands(), BinOpKind::Normal,
572 [](const APSInt &a, const APSInt &b) -> APInt { return a & b; });
573}
574
575void AndPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
576 MLIRContext *context) {
577 results
578 .insert<patterns::extendAnd, patterns::moveConstAnd, patterns::AndOfZero,
579 patterns::AndOfAllOne, patterns::AndOfSelf, patterns::AndOfPad,
580 patterns::AndOfAsSIntL, patterns::AndOfAsSIntR>(context);
581}
582
583OpFoldResult OrPrimOp::fold(FoldAdaptor adaptor) {
584 if (auto rhsCst = getConstant(adaptor.getRhs())) {
585 /// or(x, 0) -> x
586 if (rhsCst->isZero() && getLhs().getType() == getType())
587 return getLhs();
588
589 /// or(x, -1) -> -1
590 if (rhsCst->isAllOnes() && getRhs().getType() == getType() &&
591 getLhs().getType() == getType())
592 return getRhs();
593 }
594
595 if (auto lhsCst = getConstant(adaptor.getLhs())) {
596 /// or(0, x) -> x
597 if (lhsCst->isZero() && getRhs().getType() == getType())
598 return getRhs();
599
600 /// or(-1, x) -> -1
601 if (lhsCst->isAllOnes() && getLhs().getType() == getType() &&
602 getRhs().getType() == getType())
603 return getLhs();
604 }
605
606 /// or(x, x) -> x
607 if (getLhs() == getRhs() && getRhs().getType() == getType())
608 return getRhs();
609
611 *this, adaptor.getOperands(), BinOpKind::Normal,
612 [](const APSInt &a, const APSInt &b) -> APInt { return a | b; });
613}
614
615void OrPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
616 MLIRContext *context) {
617 results.insert<patterns::extendOr, patterns::moveConstOr, patterns::OrOfZero,
618 patterns::OrOfAllOne, patterns::OrOfSelf, patterns::OrOfPad,
619 patterns::OrOrr>(context);
620}
621
622OpFoldResult XorPrimOp::fold(FoldAdaptor adaptor) {
623 /// xor(x, 0) -> x
624 if (auto rhsCst = getConstant(adaptor.getRhs()))
625 if (rhsCst->isZero() &&
626 firrtl::areAnonymousTypesEquivalent(getLhs().getType(), getType()))
627 return getLhs();
628
629 /// xor(x, 0) -> x
630 if (auto lhsCst = getConstant(adaptor.getLhs()))
631 if (lhsCst->isZero() &&
632 firrtl::areAnonymousTypesEquivalent(getRhs().getType(), getType()))
633 return getRhs();
634
635 /// xor(x, x) -> 0
636 if (getLhs() == getRhs())
637 return getIntAttr(
638 getType(),
639 APInt(std::max(getType().base().getWidthOrSentinel(), 0), 0));
640
642 *this, adaptor.getOperands(), BinOpKind::Normal,
643 [](const APSInt &a, const APSInt &b) -> APInt { return a ^ b; });
644}
645
646void XorPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
647 MLIRContext *context) {
648 results.insert<patterns::extendXor, patterns::moveConstXor,
649 patterns::XorOfZero, patterns::XorOfSelf, patterns::XorOfPad>(
650 context);
651}
652
653void LEQPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
654 MLIRContext *context) {
655 results.insert<patterns::LEQWithConstLHS>(context);
656}
657
658OpFoldResult LEQPrimOp::fold(FoldAdaptor adaptor) {
659 bool isUnsigned = getLhs().getType().base().isUnsigned();
660
661 // leq(x, x) -> 1
662 if (getLhs() == getRhs())
663 return getIntAttr(getType(), APInt(1, 1));
664
665 // Comparison against constant outside type bounds.
666 if (auto width = getLhs().getType().base().getWidth()) {
667 if (auto rhsCst = getConstant(adaptor.getRhs())) {
668 auto commonWidth = std::max<int32_t>(*width, rhsCst->getBitWidth());
669 commonWidth = std::max(commonWidth, 1);
670
671 // leq(x, const) -> 0 where const < minValue of the unsigned type of x
672 // This can never occur since const is unsigned and cannot be less than 0.
673
674 // leq(x, const) -> 0 where const < minValue of the signed type of x
675 if (!isUnsigned && sextZeroWidth(*rhsCst, commonWidth)
676 .slt(getMinSignedValue(*width).sext(commonWidth)))
677 return getIntAttr(getType(), APInt(1, 0));
678
679 // leq(x, const) -> 1 where const >= maxValue of the unsigned type of x
680 if (isUnsigned && rhsCst->zext(commonWidth)
681 .uge(getMaxUnsignedValue(*width).zext(commonWidth)))
682 return getIntAttr(getType(), APInt(1, 1));
683
684 // leq(x, const) -> 1 where const >= maxValue of the signed type of x
685 if (!isUnsigned && sextZeroWidth(*rhsCst, commonWidth)
686 .sge(getMaxSignedValue(*width).sext(commonWidth)))
687 return getIntAttr(getType(), APInt(1, 1));
688 }
689 }
690
692 *this, adaptor.getOperands(), BinOpKind::Compare,
693 [=](const APSInt &a, const APSInt &b) -> APInt {
694 return APInt(1, a <= b);
695 });
696}
697
698void LTPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
699 MLIRContext *context) {
700 results.insert<patterns::LTWithConstLHS>(context);
701}
702
703OpFoldResult LTPrimOp::fold(FoldAdaptor adaptor) {
704 IntType lhsType = getLhs().getType();
705 bool isUnsigned = lhsType.isUnsigned();
706
707 // lt(x, x) -> 0
708 if (getLhs() == getRhs())
709 return getIntAttr(getType(), APInt(1, 0));
710
711 // lt(x, 0) -> 0 when x is unsigned
712 if (auto rhsCst = getConstant(adaptor.getRhs())) {
713 if (rhsCst->isZero() && lhsType.isUnsigned())
714 return getIntAttr(getType(), APInt(1, 0));
715 }
716
717 // Comparison against constant outside type bounds.
718 if (auto width = lhsType.getWidth()) {
719 if (auto rhsCst = getConstant(adaptor.getRhs())) {
720 auto commonWidth = std::max<int32_t>(*width, rhsCst->getBitWidth());
721 commonWidth = std::max(commonWidth, 1);
722
723 // lt(x, const) -> 0 where const <= minValue of the unsigned type of x
724 // Handled explicitly above.
725
726 // lt(x, const) -> 0 where const <= minValue of the signed type of x
727 if (!isUnsigned && sextZeroWidth(*rhsCst, commonWidth)
728 .sle(getMinSignedValue(*width).sext(commonWidth)))
729 return getIntAttr(getType(), APInt(1, 0));
730
731 // lt(x, const) -> 1 where const > maxValue of the unsigned type of x
732 if (isUnsigned && rhsCst->zext(commonWidth)
733 .ugt(getMaxUnsignedValue(*width).zext(commonWidth)))
734 return getIntAttr(getType(), APInt(1, 1));
735
736 // lt(x, const) -> 1 where const > maxValue of the signed type of x
737 if (!isUnsigned && sextZeroWidth(*rhsCst, commonWidth)
738 .sgt(getMaxSignedValue(*width).sext(commonWidth)))
739 return getIntAttr(getType(), APInt(1, 1));
740 }
741 }
742
744 *this, adaptor.getOperands(), BinOpKind::Compare,
745 [=](const APSInt &a, const APSInt &b) -> APInt {
746 return APInt(1, a < b);
747 });
748}
749
750void GEQPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
751 MLIRContext *context) {
752 results.insert<patterns::GEQWithConstLHS>(context);
753}
754
755OpFoldResult GEQPrimOp::fold(FoldAdaptor adaptor) {
756 IntType lhsType = getLhs().getType();
757 bool isUnsigned = lhsType.isUnsigned();
758
759 // geq(x, x) -> 1
760 if (getLhs() == getRhs())
761 return getIntAttr(getType(), APInt(1, 1));
762
763 // geq(x, 0) -> 1 when x is unsigned
764 if (auto rhsCst = getConstant(adaptor.getRhs())) {
765 if (rhsCst->isZero() && isUnsigned)
766 return getIntAttr(getType(), APInt(1, 1));
767 }
768
769 // Comparison against constant outside type bounds.
770 if (auto width = lhsType.getWidth()) {
771 if (auto rhsCst = getConstant(adaptor.getRhs())) {
772 auto commonWidth = std::max<int32_t>(*width, rhsCst->getBitWidth());
773 commonWidth = std::max(commonWidth, 1);
774
775 // geq(x, const) -> 0 where const > maxValue of the unsigned type of x
776 if (isUnsigned && rhsCst->zext(commonWidth)
777 .ugt(getMaxUnsignedValue(*width).zext(commonWidth)))
778 return getIntAttr(getType(), APInt(1, 0));
779
780 // geq(x, const) -> 0 where const > maxValue of the signed type of x
781 if (!isUnsigned && sextZeroWidth(*rhsCst, commonWidth)
782 .sgt(getMaxSignedValue(*width).sext(commonWidth)))
783 return getIntAttr(getType(), APInt(1, 0));
784
785 // geq(x, const) -> 1 where const <= minValue of the unsigned type of x
786 // Handled explicitly above.
787
788 // geq(x, const) -> 1 where const <= minValue of the signed type of x
789 if (!isUnsigned && sextZeroWidth(*rhsCst, commonWidth)
790 .sle(getMinSignedValue(*width).sext(commonWidth)))
791 return getIntAttr(getType(), APInt(1, 1));
792 }
793 }
794
796 *this, adaptor.getOperands(), BinOpKind::Compare,
797 [=](const APSInt &a, const APSInt &b) -> APInt {
798 return APInt(1, a >= b);
799 });
800}
801
802void GTPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
803 MLIRContext *context) {
804 results.insert<patterns::GTWithConstLHS>(context);
805}
806
807OpFoldResult GTPrimOp::fold(FoldAdaptor adaptor) {
808 IntType lhsType = getLhs().getType();
809 bool isUnsigned = lhsType.isUnsigned();
810
811 // gt(x, x) -> 0
812 if (getLhs() == getRhs())
813 return getIntAttr(getType(), APInt(1, 0));
814
815 // Comparison against constant outside type bounds.
816 if (auto width = lhsType.getWidth()) {
817 if (auto rhsCst = getConstant(adaptor.getRhs())) {
818 auto commonWidth = std::max<int32_t>(*width, rhsCst->getBitWidth());
819 commonWidth = std::max(commonWidth, 1);
820
821 // gt(x, const) -> 0 where const >= maxValue of the unsigned type of x
822 if (isUnsigned && rhsCst->zext(commonWidth)
823 .uge(getMaxUnsignedValue(*width).zext(commonWidth)))
824 return getIntAttr(getType(), APInt(1, 0));
825
826 // gt(x, const) -> 0 where const >= maxValue of the signed type of x
827 if (!isUnsigned && sextZeroWidth(*rhsCst, commonWidth)
828 .sge(getMaxSignedValue(*width).sext(commonWidth)))
829 return getIntAttr(getType(), APInt(1, 0));
830
831 // gt(x, const) -> 1 where const < minValue of the unsigned type of x
832 // This can never occur since const is unsigned and cannot be less than 0.
833
834 // gt(x, const) -> 1 where const < minValue of the signed type of x
835 if (!isUnsigned && sextZeroWidth(*rhsCst, commonWidth)
836 .slt(getMinSignedValue(*width).sext(commonWidth)))
837 return getIntAttr(getType(), APInt(1, 1));
838 }
839 }
840
842 *this, adaptor.getOperands(), BinOpKind::Compare,
843 [=](const APSInt &a, const APSInt &b) -> APInt {
844 return APInt(1, a > b);
845 });
846}
847
848OpFoldResult EQPrimOp::fold(FoldAdaptor adaptor) {
849 // eq(x, x) -> 1
850 if (getLhs() == getRhs())
851 return getIntAttr(getType(), APInt(1, 1));
852
853 if (auto rhsCst = getConstant(adaptor.getRhs())) {
854 /// eq(x, 1) -> x when x is 1 bit.
855 /// TODO: Support SInt<1> on the LHS etc.
856 if (rhsCst->isAllOnes() && getLhs().getType() == getType() &&
857 getRhs().getType() == getType())
858 return getLhs();
859 }
860
862 *this, adaptor.getOperands(), BinOpKind::Compare,
863 [=](const APSInt &a, const APSInt &b) -> APInt {
864 return APInt(1, a == b);
865 });
866}
867
868LogicalResult EQPrimOp::canonicalize(EQPrimOp op, PatternRewriter &rewriter) {
869 return canonicalizePrimOp(
870 op, rewriter, [&](ArrayRef<Attribute> operands) -> OpFoldResult {
871 if (auto rhsCst = getConstant(operands[1])) {
872 auto width = op.getLhs().getType().getBitWidthOrSentinel();
873
874 // eq(x, 0) -> not(x) when x is 1 bit.
875 if (rhsCst->isZero() && op.getLhs().getType() == op.getType() &&
876 op.getRhs().getType() == op.getType()) {
877 return NotPrimOp::create(rewriter, op.getLoc(), op.getLhs())
878 .getResult();
879 }
880
881 // eq(x, 0) -> not(orr(x)) when x is >1 bit
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();
885 }
886
887 // eq(x, ~0) -> andr(x) when x is >1 bit
888 if (rhsCst->isAllOnes() && width > 1 &&
889 op.getLhs().getType() == op.getRhs().getType()) {
890 return AndRPrimOp::create(rewriter, op.getLoc(), op.getLhs())
891 .getResult();
892 }
893 }
894 return {};
895 });
896}
897
898OpFoldResult NEQPrimOp::fold(FoldAdaptor adaptor) {
899 // neq(x, x) -> 0
900 if (getLhs() == getRhs())
901 return getIntAttr(getType(), APInt(1, 0));
902
903 if (auto rhsCst = getConstant(adaptor.getRhs())) {
904 /// neq(x, 0) -> x when x is 1 bit.
905 /// TODO: Support SInt<1> on the LHS etc.
906 if (rhsCst->isZero() && getLhs().getType() == getType() &&
907 getRhs().getType() == getType())
908 return getLhs();
909 }
910
912 *this, adaptor.getOperands(), BinOpKind::Compare,
913 [=](const APSInt &a, const APSInt &b) -> APInt {
914 return APInt(1, a != b);
915 });
916}
917
918LogicalResult NEQPrimOp::canonicalize(NEQPrimOp op, PatternRewriter &rewriter) {
919 return canonicalizePrimOp(
920 op, rewriter, [&](ArrayRef<Attribute> operands) -> OpFoldResult {
921 if (auto rhsCst = getConstant(operands[1])) {
922 auto width = op.getLhs().getType().getBitWidthOrSentinel();
923
924 // neq(x, 1) -> not(x) when x is 1 bit
925 if (rhsCst->isAllOnes() && op.getLhs().getType() == op.getType() &&
926 op.getRhs().getType() == op.getType()) {
927 return NotPrimOp::create(rewriter, op.getLoc(), op.getLhs())
928 .getResult();
929 }
930
931 // neq(x, 0) -> orr(x) when x is >1 bit
932 if (rhsCst->isZero() && width > 1) {
933 return OrRPrimOp::create(rewriter, op.getLoc(), op.getLhs())
934 .getResult();
935 }
936
937 // neq(x, ~0) -> not(andr(x))) when x is >1 bit
938 if (rhsCst->isAllOnes() && width > 1 &&
939 op.getLhs().getType() == op.getRhs().getType()) {
940 auto andrOp =
941 AndRPrimOp::create(rewriter, op.getLoc(), op.getLhs());
942 return NotPrimOp::create(rewriter, op.getLoc(), andrOp).getResult();
943 }
944 }
945
946 return {};
947 });
948}
949
950OpFoldResult IntegerAddOp::fold(FoldAdaptor adaptor) {
951 // TODO: implement constant folding, etc.
952 // Tracked in https://github.com/llvm/circt/issues/6696.
953 return {};
954}
955
956OpFoldResult IntegerMulOp::fold(FoldAdaptor adaptor) {
957 // TODO: implement constant folding, etc.
958 // Tracked in https://github.com/llvm/circt/issues/6724.
959 return {};
960}
961
962OpFoldResult IntegerShrOp::fold(FoldAdaptor adaptor) {
963 if (auto rhsCst = getConstant(adaptor.getRhs())) {
964 if (auto lhsCst = getConstant(adaptor.getLhs())) {
965
966 return IntegerAttr::get(IntegerType::get(getContext(),
967 lhsCst->getBitWidth(),
968 IntegerType::Signed),
969 lhsCst->ashr(*rhsCst));
970 }
971
972 if (rhsCst->isZero())
973 return getLhs();
974 }
975
976 return {};
977}
978
979OpFoldResult IntegerShlOp::fold(FoldAdaptor adaptor) {
980 if (auto rhsCst = getConstant(adaptor.getRhs())) {
981 // Constant folding
982 if (auto lhsCst = getConstant(adaptor.getLhs()))
983
984 return IntegerAttr::get(IntegerType::get(getContext(),
985 lhsCst->getBitWidth(),
986 IntegerType::Signed),
987 lhsCst->shl(*rhsCst));
988
989 // integer.shl(x, 0) -> x
990 if (rhsCst->isZero())
991 return getLhs();
992 }
993
994 return {};
995}
996
997//===----------------------------------------------------------------------===//
998// Unary Operators
999//===----------------------------------------------------------------------===//
1000
1001OpFoldResult SizeOfIntrinsicOp::fold(FoldAdaptor) {
1002 auto base = getInput().getType();
1003 auto w = getBitWidth(base);
1004 if (w)
1005 return getIntAttr(getType(), APInt(32, *w));
1006 return {};
1007}
1008
1009OpFoldResult IsXIntrinsicOp::fold(FoldAdaptor adaptor) {
1010 // No constant can be 'x' by definition.
1011 if (auto cst = getConstant(adaptor.getArg()))
1012 return getIntAttr(getType(), APInt(1, 0));
1013 return {};
1014}
1015
1016OpFoldResult AsSIntPrimOp::fold(FoldAdaptor adaptor) {
1017 // No effect.
1018 if (areAnonymousTypesEquivalent(getInput().getType(), getType()))
1019 return getInput();
1020
1021 // Be careful to only fold the cast into the constant if the size is known.
1022 // Otherwise width inference may produce differently-sized constants if the
1023 // sign changes.
1024 if (getType().base().hasWidth())
1025 if (auto cst = getConstant(adaptor.getInput()))
1026 return getIntAttr(getType(), *cst);
1027
1028 return {};
1029}
1030
1031void AsSIntPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1032 MLIRContext *context) {
1033 results.insert<patterns::StoUtoS>(context);
1034}
1035
1036OpFoldResult AsUIntPrimOp::fold(FoldAdaptor adaptor) {
1037 // No effect.
1038 if (areAnonymousTypesEquivalent(getInput().getType(), getType()))
1039 return getInput();
1040
1041 // Be careful to only fold the cast into the constant if the size is known.
1042 // Otherwise width inference may produce differently-sized constants if the
1043 // sign changes.
1044 if (getType().base().hasWidth())
1045 if (auto cst = getConstant(adaptor.getInput()))
1046 return getIntAttr(getType(), *cst);
1047
1048 return {};
1049}
1050
1051void AsUIntPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1052 MLIRContext *context) {
1053 results.insert<patterns::UtoStoU>(context);
1054}
1055
1056OpFoldResult AsAsyncResetPrimOp::fold(FoldAdaptor adaptor) {
1057 // No effect.
1058 if (getInput().getType() == getType())
1059 return getInput();
1060
1061 // Constant fold.
1062 if (auto cst = getConstant(adaptor.getInput()))
1063 return BoolAttr::get(getContext(), cst->getBoolValue());
1064
1065 return {};
1066}
1067
1068OpFoldResult AsResetPrimOp::fold(FoldAdaptor adaptor) {
1069 if (auto cst = getConstant(adaptor.getInput()))
1070 return BoolAttr::get(getContext(), cst->getBoolValue());
1071 return {};
1072}
1073
1074OpFoldResult AsClockPrimOp::fold(FoldAdaptor adaptor) {
1075 // No effect.
1076 if (getInput().getType() == getType())
1077 return getInput();
1078
1079 // Constant fold.
1080 if (auto cst = getConstant(adaptor.getInput()))
1081 return BoolAttr::get(getContext(), cst->getBoolValue());
1082
1083 return {};
1084}
1085
1086OpFoldResult CvtPrimOp::fold(FoldAdaptor adaptor) {
1087 if (!hasKnownWidthIntTypes(*this))
1088 return {};
1089
1090 // Signed to signed is a noop, unsigned operands prepend a zero bit.
1091 if (auto cst = getExtendedConstant(getOperand(), adaptor.getInput(),
1092 getType().base().getWidthOrSentinel()))
1093 return getIntAttr(getType(), *cst);
1094
1095 return {};
1096}
1097
1098void CvtPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1099 MLIRContext *context) {
1100 results.insert<patterns::CVTSigned, patterns::CVTUnSigned>(context);
1101}
1102
1103OpFoldResult NegPrimOp::fold(FoldAdaptor adaptor) {
1104 if (!hasKnownWidthIntTypes(*this))
1105 return {};
1106
1107 // FIRRTL negate always adds a bit.
1108 // -x ---> 0-sext(x) or 0-zext(x)
1109 if (auto cst = getExtendedConstant(getOperand(), adaptor.getInput(),
1110 getType().base().getWidthOrSentinel()))
1111 return getIntAttr(getType(), APInt((*cst).getBitWidth(), 0) - *cst);
1112
1113 return {};
1114}
1115
1116OpFoldResult NotPrimOp::fold(FoldAdaptor adaptor) {
1117 if (!hasKnownWidthIntTypes(*this))
1118 return {};
1119
1120 if (auto cst = getExtendedConstant(getOperand(), adaptor.getInput(),
1121 getType().base().getWidthOrSentinel()))
1122 return getIntAttr(getType(), ~*cst);
1123
1124 return {};
1125}
1126
1127void NotPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1128 MLIRContext *context) {
1129 results.insert<patterns::NotNot, patterns::NotEq, patterns::NotNeq,
1130 patterns::NotLeq, patterns::NotLt, patterns::NotGeq,
1131 patterns::NotGt>(context);
1132}
1133
1134class ReductionCat : public mlir::RewritePattern {
1135public:
1136 ReductionCat(MLIRContext *context, llvm::StringLiteral opName)
1137 : RewritePattern(opName, 0, context) {}
1138
1139 /// Handle a constant operand in the cat operation.
1140 /// Returns true if the entire reduction can be replaced with a constant.
1141 /// May add non-zero constants to the remaining operands list.
1142 virtual bool handleConstant(mlir::PatternRewriter &rewriter, Operation *op,
1143 ConstantOp constantOp,
1144 SmallVectorImpl<Value> &remaining) const = 0;
1145
1146 /// Return the unit value for this reduction operation:
1147 virtual bool getIdentityValue() const = 0;
1148
1149 LogicalResult
1150 matchAndRewrite(Operation *op,
1151 mlir::PatternRewriter &rewriter) const override {
1152 // Check if the operand is a cat operation
1153 auto catOp = op->getOperand(0).getDefiningOp<CatPrimOp>();
1154 if (!catOp)
1155 return failure();
1156
1157 SmallVector<Value> nonConstantOperands;
1158
1159 // Process each operand of the cat operation
1160 for (auto operand : catOp.getInputs()) {
1161 if (auto constantOp = operand.getDefiningOp<ConstantOp>()) {
1162 // Handle constant operands - may short-circuit the entire operation
1163 if (handleConstant(rewriter, op, constantOp, nonConstantOperands))
1164 return success();
1165 } else {
1166 // Keep non-constant operands for further processing
1167 nonConstantOperands.push_back(operand);
1168 }
1169 }
1170
1171 // If no operands remain, replace with identity value
1172 if (nonConstantOperands.empty()) {
1173 replaceOpWithNewOpAndCopyName<ConstantOp>(
1174 rewriter, op, cast<IntType>(op->getResult(0).getType()),
1175 APInt(1, getIdentityValue()));
1176 return success();
1177 }
1178
1179 // If only one operand remains, apply reduction directly to it
1180 if (nonConstantOperands.size() == 1) {
1181 rewriter.modifyOpInPlace(
1182 op, [&] { op->setOperand(0, nonConstantOperands.front()); });
1183 return success();
1184 }
1185
1186 // Multiple operands remain - optimize only when cat has a single use.
1187 if (catOp->hasOneUse() &&
1188 nonConstantOperands.size() < catOp->getNumOperands()) {
1189 replaceOpWithNewOpAndCopyName<CatPrimOp>(rewriter, catOp,
1190 nonConstantOperands);
1191 return success();
1192 }
1193 return failure();
1194 }
1195};
1196
1197class OrRCat : public ReductionCat {
1198public:
1199 OrRCat(MLIRContext *context)
1200 : ReductionCat(context, OrRPrimOp::getOperationName()) {}
1201 bool handleConstant(mlir::PatternRewriter &rewriter, Operation *op,
1202 ConstantOp value,
1203 SmallVectorImpl<Value> &remaining) const override {
1204 if (value.getValue().isZero())
1205 return false;
1206
1207 replaceOpWithNewOpAndCopyName<ConstantOp>(
1208 rewriter, op, cast<IntType>(op->getResult(0).getType()),
1209 llvm::APInt(1, 1));
1210 return true;
1211 }
1212 bool getIdentityValue() const override { return false; }
1213};
1214
1215class AndRCat : public ReductionCat {
1216public:
1217 AndRCat(MLIRContext *context)
1218 : ReductionCat(context, AndRPrimOp::getOperationName()) {}
1219 bool handleConstant(mlir::PatternRewriter &rewriter, Operation *op,
1220 ConstantOp value,
1221 SmallVectorImpl<Value> &remaining) const override {
1222 if (value.getValue().isAllOnes())
1223 return false;
1224
1225 replaceOpWithNewOpAndCopyName<ConstantOp>(
1226 rewriter, op, cast<IntType>(op->getResult(0).getType()),
1227 llvm::APInt(1, 0));
1228 return true;
1229 }
1230 bool getIdentityValue() const override { return true; }
1231};
1232
1233class XorRCat : public ReductionCat {
1234public:
1235 XorRCat(MLIRContext *context)
1236 : ReductionCat(context, XorRPrimOp::getOperationName()) {}
1237 bool handleConstant(mlir::PatternRewriter &rewriter, Operation *op,
1238 ConstantOp value,
1239 SmallVectorImpl<Value> &remaining) const override {
1240 if (value.getValue().isZero())
1241 return false;
1242 remaining.push_back(value);
1243 return false;
1244 }
1245 bool getIdentityValue() const override { return false; }
1246};
1247
1248OpFoldResult AndRPrimOp::fold(FoldAdaptor adaptor) {
1249 if (!hasKnownWidthIntTypes(*this))
1250 return {};
1251
1252 if (getInput().getType().getBitWidthOrSentinel() == 0)
1253 return getIntAttr(getType(), APInt(1, 1));
1254
1255 // x == -1
1256 if (auto cst = getConstant(adaptor.getInput()))
1257 return getIntAttr(getType(), APInt(1, cst->isAllOnes()));
1258
1259 // one bit is identity. Only applies to UInt since we can't make a cast
1260 // here.
1261 if (isUInt1(getInput().getType()))
1262 return getInput();
1263
1264 return {};
1265}
1266
1267void AndRPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1268 MLIRContext *context) {
1269 results.insert<patterns::AndRasSInt, patterns::AndRasUInt, patterns::AndRPadU,
1270 patterns::AndRPadS, patterns::AndRCatAndR_left,
1271 patterns::AndRCatAndR_right, AndRCat>(context);
1272}
1273
1274OpFoldResult OrRPrimOp::fold(FoldAdaptor adaptor) {
1275 if (!hasKnownWidthIntTypes(*this))
1276 return {};
1277
1278 if (getInput().getType().getBitWidthOrSentinel() == 0)
1279 return getIntAttr(getType(), APInt(1, 0));
1280
1281 // x != 0
1282 if (auto cst = getConstant(adaptor.getInput()))
1283 return getIntAttr(getType(), APInt(1, !cst->isZero()));
1284
1285 // one bit is identity. Only applies to UInt since we can't make a cast
1286 // here.
1287 if (isUInt1(getInput().getType()))
1288 return getInput();
1289
1290 return {};
1291}
1292
1293void OrRPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1294 MLIRContext *context) {
1295 results.insert<patterns::OrRasSInt, patterns::OrRasUInt, patterns::OrRPadU,
1296 patterns::OrRCatOrR_left, patterns::OrRCatOrR_right, OrRCat>(
1297 context);
1298}
1299
1300OpFoldResult XorRPrimOp::fold(FoldAdaptor adaptor) {
1301 if (!hasKnownWidthIntTypes(*this))
1302 return {};
1303
1304 if (getInput().getType().getBitWidthOrSentinel() == 0)
1305 return getIntAttr(getType(), APInt(1, 0));
1306
1307 // popcount(x) & 1
1308 if (auto cst = getConstant(adaptor.getInput()))
1309 return getIntAttr(getType(), APInt(1, cst->popcount() & 1));
1310
1311 // one bit is identity. Only applies to UInt since we can't make a cast here.
1312 if (isUInt1(getInput().getType()))
1313 return getInput();
1314
1315 return {};
1316}
1317
1318void XorRPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1319 MLIRContext *context) {
1320 results
1321 .insert<patterns::XorRasSInt, patterns::XorRasUInt, patterns::XorRPadU,
1322 patterns::XorRCatXorR_left, patterns::XorRCatXorR_right, XorRCat>(
1323 context);
1324}
1325
1326//===----------------------------------------------------------------------===//
1327// Other Operators
1328//===----------------------------------------------------------------------===//
1329
1330OpFoldResult CatPrimOp::fold(FoldAdaptor adaptor) {
1331 auto inputs = getInputs();
1332 auto inputAdaptors = adaptor.getInputs();
1333
1334 // If no inputs, return 0-bit value
1335 if (inputs.empty())
1336 return getIntZerosAttr(getType());
1337
1338 // If single input and same type, return it
1339 if (inputs.size() == 1 && inputs[0].getType() == getType())
1340 return inputs[0];
1341
1342 // Make sure it's safe to fold.
1343 if (!hasKnownWidthIntTypes(*this))
1344 return {};
1345
1346 // Filter out zero-width operands
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);
1354 if (!attr)
1355 allConstant = false;
1356 if (nonZeroInputs.size() > 1 && !allConstant)
1357 return {};
1358 }
1359 }
1360
1361 // If all inputs were zero-width, return 0-bit value
1362 if (nonZeroInputs.empty())
1363 return getIntZerosAttr(getType());
1364
1365 // If only one non-zero input and it has the same type as result, return it
1366 if (nonZeroInputs.size() == 1 && nonZeroInputs[0].getType() == getType())
1367 return nonZeroInputs[0];
1368
1369 if (!hasKnownWidthIntTypes(*this))
1370 return {};
1371
1372 // Constant fold cat - concatenate all constant operands
1373 SmallVector<APInt> constants;
1374 for (auto inputAdaptor : inputAdaptors) {
1375 if (auto cst = getConstant(inputAdaptor))
1376 constants.push_back(*cst);
1377 else
1378 return {}; // Not all operands are constant
1379 }
1380
1381 assert(!constants.empty());
1382 // Concatenate all constants from left to right
1383 APInt result = constants[0];
1384 for (size_t i = 1; i < constants.size(); ++i)
1385 result = result.concat(constants[i]);
1386
1387 return getIntAttr(getType(), result);
1388}
1389
1390void DShlPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1391 MLIRContext *context) {
1392 results.insert<patterns::DShlOfConstant>(context);
1393}
1394
1395void DShrPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1396 MLIRContext *context) {
1397 results.insert<patterns::DShrOfConstant>(context);
1398}
1399
1400namespace {
1401
1402/// Canonicalize a cat tree into single cat operation.
1403class FlattenCat : public mlir::OpRewritePattern<CatPrimOp> {
1404public:
1405 using OpRewritePattern::OpRewritePattern;
1406
1407 LogicalResult
1408 matchAndRewrite(CatPrimOp cat,
1409 mlir::PatternRewriter &rewriter) const override {
1410 if (!hasKnownWidthIntTypes(cat) ||
1411 cat.getType().getBitWidthOrSentinel() == 0)
1412 return failure();
1413
1414 // If this is not a root of concat tree, skip.
1415 if (cat->hasOneUse() && isa<CatPrimOp>(*cat->getUsers().begin()))
1416 return failure();
1417
1418 // Try flattening the cat tree.
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);
1424 };
1425 pushOperands(cat);
1426 bool hasSigned = false, hasUnsigned = false;
1427 while (!worklist.empty()) {
1428 auto value = worklist.pop_back_val();
1429 auto catOp = value.getDefiningOp<CatPrimOp>();
1430 if (!catOp) {
1431 operands.push_back(value);
1432 (type_isa<UIntType>(value.getType()) ? hasUnsigned : hasSigned) = true;
1433 continue;
1434 }
1435
1436 pushOperands(catOp);
1437 }
1438
1439 // Helper function to cast signed values to unsigned. CatPrimOp converts
1440 // signed values to unsigned so we need to do it explicitly here.
1441 auto castToUIntIfSigned = [&](Value value) -> Value {
1442 if (type_isa<UIntType>(value.getType()))
1443 return value;
1444 return AsUIntPrimOp::create(rewriter, value.getLoc(), value);
1445 };
1446
1447 assert(operands.size() >= 1 && "zero width cast must be rejected");
1448
1449 if (operands.size() == 1) {
1450 rewriter.replaceOp(cat, castToUIntIfSigned(operands[0]));
1451 return success();
1452 }
1453
1454 if (operands.size() == cat->getNumOperands())
1455 return failure();
1456
1457 // If types are mixed, cast all operands to unsigned.
1458 if (hasSigned && hasUnsigned)
1459 for (auto &operand : operands)
1460 operand = castToUIntIfSigned(operand);
1461
1462 replaceOpWithNewOpAndCopyName<CatPrimOp>(rewriter, cat, cat.getType(),
1463 operands);
1464 return success();
1465 }
1466};
1467
1468// Fold successive constants cat(x, y, 1, 10, z) -> cat(x, y, 110, z)
1469class CatOfConstant : public mlir::OpRewritePattern<CatPrimOp> {
1470public:
1471 using OpRewritePattern::OpRewritePattern;
1472
1473 LogicalResult
1474 matchAndRewrite(CatPrimOp cat,
1475 mlir::PatternRewriter &rewriter) const override {
1476 if (!hasKnownWidthIntTypes(cat))
1477 return failure();
1478
1479 SmallVector<Value> operands;
1480
1481 for (size_t i = 0; i < cat->getNumOperands(); ++i) {
1482 auto cst = cat.getInputs()[i].getDefiningOp<ConstantOp>();
1483 if (!cst) {
1484 operands.push_back(cat.getInputs()[i]);
1485 continue;
1486 }
1487 APSInt value = cst.getValue();
1488 size_t j = i + 1;
1489 for (; j < cat->getNumOperands(); ++j) {
1490 auto nextCst = cat.getInputs()[j].getDefiningOp<ConstantOp>();
1491 if (!nextCst)
1492 break;
1493 value = value.concat(nextCst.getValue());
1494 }
1495
1496 if (j == i + 1) {
1497 // Not folded.
1498 operands.push_back(cst);
1499 } else {
1500 // Folded.
1501 operands.push_back(ConstantOp::create(rewriter, cat.getLoc(), value));
1502 }
1503
1504 i = j - 1;
1505 }
1506
1507 if (operands.size() == cat->getNumOperands())
1508 return failure();
1509
1510 replaceOpWithNewOpAndCopyName<CatPrimOp>(rewriter, cat, cat.getType(),
1511 operands);
1512
1513 return success();
1514 }
1515};
1516
1517} // namespace
1518
1519void CatPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1520 MLIRContext *context) {
1521 results.insert<patterns::CatBitsBits, patterns::CatDoubleConst,
1522 patterns::CatCast, FlattenCat, CatOfConstant>(context);
1523}
1524
1525//===----------------------------------------------------------------------===//
1526// StringConcatOp
1527//===----------------------------------------------------------------------===//
1528
1529OpFoldResult StringConcatOp::fold(FoldAdaptor adaptor) {
1530 // Fold single-operand concat to just the operand.
1531 if (getInputs().size() == 1)
1532 return getInputs()[0];
1533
1534 // Check if all operands are constant strings before accumulating.
1535 if (!llvm::all_of(adaptor.getInputs(), [](Attribute operand) {
1536 return isa_and_nonnull<StringAttr>(operand);
1537 }))
1538 return {};
1539
1540 // All operands are constant strings, concatenate them.
1541 SmallString<64> result;
1542 for (auto operand : adaptor.getInputs())
1543 result += cast<StringAttr>(operand).getValue();
1544
1545 return StringAttr::get(getContext(), result);
1546}
1547
1548namespace {
1549/// Flatten nested string.concat operations into a single concat.
1550/// string.concat(a, string.concat(b, c), d) -> string.concat(a, b, c, d)
1551class FlattenStringConcat : public mlir::OpRewritePattern<StringConcatOp> {
1552public:
1553 using OpRewritePattern::OpRewritePattern;
1554
1555 LogicalResult
1556 matchAndRewrite(StringConcatOp concat,
1557 mlir::PatternRewriter &rewriter) const override {
1558
1559 // Check if any operands are nested concats with a single use. Only inline
1560 // single-use nested concats to avoid fighting with DCE.
1561 bool hasNestedConcat = llvm::any_of(concat.getInputs(), [](Value operand) {
1562 auto nestedConcat = operand.getDefiningOp<StringConcatOp>();
1563 return nestedConcat && operand.hasOneUse();
1564 });
1565
1566 if (!hasNestedConcat)
1567 return failure();
1568
1569 // Flatten nested concats that have a single use.
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());
1575 else
1576 flatOperands.push_back(input);
1577 }
1578
1579 rewriter.modifyOpInPlace(concat,
1580 [&]() { concat->setOperands(flatOperands); });
1581 return success();
1582 }
1583};
1584
1585/// Merge consecutive constant strings in a concat and remove empty strings.
1586/// string.concat("a", "b", x, "", "c", "d") -> string.concat("ab", x, "cd")
1587class MergeAdjacentStringConstants
1588 : public mlir::OpRewritePattern<StringConcatOp> {
1589public:
1590 using OpRewritePattern::OpRewritePattern;
1591
1592 LogicalResult
1593 matchAndRewrite(StringConcatOp concat,
1594 mlir::PatternRewriter &rewriter) const override {
1595
1596 SmallVector<Value> newOperands;
1597 SmallString<64> accumulatedLit;
1598 SmallVector<StringConstantOp> accumulatedOps;
1599 bool changed = false;
1600
1601 auto flushLiterals = [&]() {
1602 if (accumulatedOps.empty())
1603 return;
1604
1605 // If only one literal, reuse it.
1606 if (accumulatedOps.size() == 1) {
1607 newOperands.push_back(accumulatedOps[0]);
1608 } else {
1609 // Multiple literals - merge them.
1610 auto newLit = rewriter.createOrFold<StringConstantOp>(
1611 concat.getLoc(), StringAttr::get(getContext(), accumulatedLit));
1612 newOperands.push_back(newLit);
1613 changed = true;
1614 }
1615 accumulatedLit.clear();
1616 accumulatedOps.clear();
1617 };
1618
1619 for (auto operand : concat.getInputs()) {
1620 if (auto litOp = operand.getDefiningOp<StringConstantOp>()) {
1621 // Skip empty strings.
1622 if (litOp.getValue().empty()) {
1623 changed = true;
1624 continue;
1625 }
1626 accumulatedLit += litOp.getValue();
1627 accumulatedOps.push_back(litOp);
1628 } else {
1629 flushLiterals();
1630 newOperands.push_back(operand);
1631 }
1632 }
1633
1634 // Flush any remaining literals.
1635 flushLiterals();
1636
1637 if (!changed)
1638 return failure();
1639
1640 // If no operands remain, replace with empty string.
1641 if (newOperands.empty())
1642 return rewriter.replaceOpWithNewOp<StringConstantOp>(
1643 concat, StringAttr::get(getContext(), "")),
1644 success();
1645
1646 // Single-operand case is handled by the folder.
1647 rewriter.modifyOpInPlace(concat,
1648 [&]() { concat->setOperands(newOperands); });
1649 return success();
1650 }
1651};
1652
1653} // namespace
1654
1655void StringConcatOp::getCanonicalizationPatterns(RewritePatternSet &results,
1656 MLIRContext *context) {
1657 results.insert<FlattenStringConcat, MergeAdjacentStringConstants>(context);
1658}
1659
1660//===----------------------------------------------------------------------===//
1661// PropEqOp
1662//===----------------------------------------------------------------------===//
1663
1664OpFoldResult PropEqOp::fold(FoldAdaptor adaptor) {
1665 auto lhsAttr = adaptor.getLhs();
1666 auto rhsAttr = adaptor.getRhs();
1667 if (!lhsAttr || !rhsAttr)
1668 return {};
1669
1670 return BoolAttr::get(getContext(), lhsAttr == rhsAttr);
1671}
1672
1673//===----------------------------------------------------------------------===//
1674// Boolean Property Ops
1675//===----------------------------------------------------------------------===//
1676
1677// Helper to extract the bool value from a BoolAttr.
1678static std::optional<bool> getBoolValue(Attribute attr) {
1679 if (auto boolAttr = dyn_cast_or_null<BoolAttr>(attr))
1680 return boolAttr.getValue();
1681 return std::nullopt;
1682}
1683
1684OpFoldResult BoolAndOp::fold(FoldAdaptor adaptor) {
1685 auto lhs = getBoolValue(adaptor.getLhs());
1686 auto rhs = getBoolValue(adaptor.getRhs());
1687 if (lhs && rhs)
1688 return BoolAttr::get(getContext(), *lhs && *rhs);
1689 // AND with false is always false.
1690 if ((lhs && !*lhs) || (rhs && !*rhs))
1691 return BoolAttr::get(getContext(), false);
1692 // AND with true is identity.
1693 if (lhs && *lhs)
1694 return getRhs();
1695 if (rhs && *rhs)
1696 return getLhs();
1697 return {};
1698}
1699
1700OpFoldResult BoolOrOp::fold(FoldAdaptor adaptor) {
1701 auto lhs = getBoolValue(adaptor.getLhs());
1702 auto rhs = getBoolValue(adaptor.getRhs());
1703 if (lhs && rhs)
1704 return BoolAttr::get(getContext(), *lhs || *rhs);
1705 // OR with true is always true.
1706 if ((lhs && *lhs) || (rhs && *rhs))
1707 return BoolAttr::get(getContext(), true);
1708 // OR with false is identity.
1709 if (lhs && !*lhs)
1710 return getRhs();
1711 if (rhs && !*rhs)
1712 return getLhs();
1713 return {};
1714}
1715
1716OpFoldResult BoolXorOp::fold(FoldAdaptor adaptor) {
1717 auto lhs = getBoolValue(adaptor.getLhs());
1718 auto rhs = getBoolValue(adaptor.getRhs());
1719 if (lhs && rhs)
1720 return BoolAttr::get(getContext(), *lhs ^ *rhs);
1721 // XOR with false is identity.
1722 if (lhs && !*lhs)
1723 return getRhs();
1724 if (rhs && !*rhs)
1725 return getLhs();
1726 return {};
1727}
1728
1729OpFoldResult BitCastOp::fold(FoldAdaptor adaptor) {
1730 auto op = (*this);
1731 // BitCast is redundant if input and result types are same.
1732 if (op.getType() == op.getInput().getType())
1733 return op.getInput();
1734
1735 // Two consecutive BitCasts are redundant if first bitcast type is same as the
1736 // final result type.
1737 if (BitCastOp in = dyn_cast_or_null<BitCastOp>(op.getInput().getDefiningOp()))
1738 if (op.getType() == in.getInput().getType())
1739 return in.getInput();
1740
1741 return {};
1742}
1743
1744OpFoldResult BitsPrimOp::fold(FoldAdaptor adaptor) {
1745 IntType inputType = getInput().getType();
1746 IntType resultType = getType();
1747 // If we are extracting the entire input, then return it.
1748 if (inputType == getType() && resultType.hasWidth())
1749 return getInput();
1750
1751 // Constant fold.
1752 if (hasKnownWidthIntTypes(*this))
1753 if (auto cst = getConstant(adaptor.getInput()))
1754 return getIntAttr(resultType,
1755 cst->extractBits(getHi() - getLo() + 1, getLo()));
1756
1757 return {};
1758}
1759
1760struct BitsOfCat : public mlir::OpRewritePattern<BitsPrimOp> {
1761 using OpRewritePattern::OpRewritePattern;
1762
1763 LogicalResult
1764 matchAndRewrite(BitsPrimOp bits,
1765 mlir::PatternRewriter &rewriter) const override {
1766 auto cat = bits.getInput().getDefiningOp<CatPrimOp>();
1767 if (!cat)
1768 return failure();
1769 int32_t bitPos = bits.getLo();
1770 auto resultWidth = type_cast<UIntType>(bits.getType()).getWidthOrSentinel();
1771 if (resultWidth < 0)
1772 return failure();
1773 for (auto operand : llvm::reverse(cat.getInputs())) {
1774 auto operandWidth =
1775 type_cast<IntType>(operand.getType()).getWidthOrSentinel();
1776 if (operandWidth < 0)
1777 return failure();
1778 if (bitPos < operandWidth) {
1779 if (bitPos + resultWidth <= operandWidth) {
1780 auto newBits = rewriter.createOrFold<BitsPrimOp>(
1781 bits.getLoc(), operand, bitPos + resultWidth - 1, bitPos);
1782 replaceOpAndCopyName(rewriter, bits, newBits);
1783 return success();
1784 }
1785 return failure();
1786 }
1787 bitPos -= operandWidth;
1788 }
1789 return failure();
1790 }
1791};
1792
1793void BitsPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1794 MLIRContext *context) {
1795 results
1796 .insert<patterns::BitsOfBits, patterns::BitsOfMux, patterns::BitsOfAsUInt,
1797 patterns::BitsOfAnd, patterns::BitsOfPad, BitsOfCat>(context);
1798}
1799
1800/// Replace the specified operation with a 'bits' op from the specified hi/lo
1801/// bits. Insert a cast to handle the case where the original operation
1802/// returned a signed integer.
1803static void replaceWithBits(Operation *op, Value value, unsigned hiBit,
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);
1808
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);
1814 }
1815 rewriter.replaceOp(op, value);
1816}
1817
1818template <typename OpTy>
1819static OpFoldResult foldMux(OpTy op, typename OpTy::FoldAdaptor adaptor) {
1820 // mux : UInt<0> -> 0
1821 if (op.getType().getBitWidthOrSentinel() == 0)
1822 return getIntAttr(op.getType(),
1823 APInt(0, 0, op.getType().isSignedInteger()));
1824
1825 // mux(cond, x, x) -> x
1826 if (op.getHigh() == op.getLow() && op.getHigh().getType() == op.getType())
1827 return op.getHigh();
1828
1829 // The following folds require that the result has a known width. Otherwise
1830 // the mux requires an additional padding operation to be inserted, which is
1831 // not possible in a fold.
1832 if (op.getType().getBitWidthOrSentinel() < 0)
1833 return {};
1834
1835 // mux(0/1, x, y) -> x or y
1836 if (auto cond = getConstant(adaptor.getSel())) {
1837 if (cond->isZero() && op.getLow().getType() == op.getType())
1838 return op.getLow();
1839 if (!cond->isZero() && op.getHigh().getType() == op.getType())
1840 return op.getHigh();
1841 }
1842
1843 // mux(cond, x, cst)
1844 if (auto lowCst = getConstant(adaptor.getLow())) {
1845 // mux(cond, c1, c2)
1846 if (auto highCst = getConstant(adaptor.getHigh())) {
1847 // mux(cond, cst, cst) -> cst
1848 if (highCst->getBitWidth() == lowCst->getBitWidth() &&
1849 *highCst == *lowCst)
1850 // Ensure that this has an integer representation. This specifically
1851 // avoids problems with clocks.
1852 if (auto intType = type_dyn_cast<IntType>(op.getType()))
1853 if (intType.hasWidth() &&
1854 (unsigned)intType.getWidthOrSentinel() == highCst->getBitWidth())
1855 return getIntAttr(op.getType(), *highCst);
1856 // mux(cond, 1, 0) -> cond
1857 if (highCst->isOne() && lowCst->isZero() &&
1858 op.getType() == op.getSel().getType())
1859 return op.getSel();
1860
1861 // TODO: x ? ~0 : 0 -> sext(x)
1862 // TODO: "x ? c1 : c2" -> many tricks
1863 }
1864 // TODO: "x ? a : 0" -> sext(x) & a
1865 }
1866
1867 // TODO: "x ? c1 : y" -> "~x ? y : c1"
1868 return {};
1869}
1870
1871OpFoldResult MuxPrimOp::fold(FoldAdaptor adaptor) {
1872 return foldMux(*this, adaptor);
1873}
1874
1875OpFoldResult Mux2CellIntrinsicOp::fold(FoldAdaptor adaptor) {
1876 return foldMux(*this, adaptor);
1877}
1878
1879OpFoldResult Mux4CellIntrinsicOp::fold(FoldAdaptor adaptor) { return {}; }
1880
1881namespace {
1882
1883// If the mux has a known output width, pad the operands up to this width.
1884// Most folds on mux require that folded operands are of the same width as
1885// the mux itself.
1886class MuxPad : public mlir::OpRewritePattern<MuxPrimOp> {
1887public:
1888 using OpRewritePattern::OpRewritePattern;
1889
1890 LogicalResult
1891 matchAndRewrite(MuxPrimOp mux,
1892 mlir::PatternRewriter &rewriter) const override {
1893 auto width = mux.getType().getBitWidthOrSentinel();
1894 if (width < 0)
1895 return failure();
1896
1897 auto pad = [&](Value input) -> Value {
1898 auto inputWidth =
1899 type_cast<FIRRTLBaseType>(input.getType()).getBitWidthOrSentinel();
1900 if (inputWidth < 0 || width == inputWidth)
1901 return input;
1902 return PadPrimOp::create(rewriter, mux.getLoc(), mux.getType(), input,
1903 width)
1904 .getResult();
1905 };
1906
1907 auto newHigh = pad(mux.getHigh());
1908 auto newLow = pad(mux.getLow());
1909 if (newHigh == mux.getHigh() && newLow == mux.getLow())
1910 return failure();
1911
1912 replaceOpWithNewOpAndCopyName<MuxPrimOp>(
1913 rewriter, mux, mux.getType(), ValueRange{mux.getSel(), newHigh, newLow},
1914 mux->getAttrs());
1915 return success();
1916 }
1917};
1918
1919// Find muxes which have conditions dominated by other muxes with the same
1920// condition.
1921class MuxSharedCond : public mlir::OpRewritePattern<MuxPrimOp> {
1922public:
1923 using OpRewritePattern::OpRewritePattern;
1924
1925 static const int depthLimit = 5;
1926
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);
1934 });
1935 return {};
1936 }
1937 rewriter.setInsertionPointAfter(mux);
1938 return MuxPrimOp::create(rewriter, mux.getLoc(), mux.getType(),
1939 ValueRange{mux.getSel(), high, low})
1940 .getResult();
1941 }
1942
1943 // Walk a dependent mux tree assuming the condition cond is true.
1944 Value tryCondTrue(Value op, Value cond, mlir::PatternRewriter &rewriter,
1945 bool updateInPlace, int limit) const {
1946 MuxPrimOp mux = op.getDefiningOp<MuxPrimOp>();
1947 if (!mux)
1948 return {};
1949 if (mux.getSel() == cond)
1950 return mux.getHigh();
1951 if (limit > depthLimit)
1952 return {};
1953 updateInPlace &= mux->hasOneUse();
1954
1955 if (Value v = tryCondTrue(mux.getHigh(), cond, rewriter, updateInPlace,
1956 limit + 1))
1957 return updateOrClone(mux, v, mux.getLow(), rewriter, updateInPlace);
1958
1959 if (Value v =
1960 tryCondTrue(mux.getLow(), cond, rewriter, updateInPlace, limit + 1))
1961 return updateOrClone(mux, mux.getHigh(), v, rewriter, updateInPlace);
1962 return {};
1963 }
1964
1965 // Walk a dependent mux tree assuming the condition cond is false.
1966 Value tryCondFalse(Value op, Value cond, mlir::PatternRewriter &rewriter,
1967 bool updateInPlace, int limit) const {
1968 MuxPrimOp mux = op.getDefiningOp<MuxPrimOp>();
1969 if (!mux)
1970 return {};
1971 if (mux.getSel() == cond)
1972 return mux.getLow();
1973 if (limit > depthLimit)
1974 return {};
1975 updateInPlace &= mux->hasOneUse();
1976
1977 if (Value v = tryCondFalse(mux.getHigh(), cond, rewriter, updateInPlace,
1978 limit + 1))
1979 return updateOrClone(mux, v, mux.getLow(), rewriter, updateInPlace);
1980
1981 if (Value v = tryCondFalse(mux.getLow(), cond, rewriter, updateInPlace,
1982 limit + 1))
1983 return updateOrClone(mux, mux.getHigh(), v, rewriter, updateInPlace);
1984
1985 return {};
1986 }
1987
1988 LogicalResult
1989 matchAndRewrite(MuxPrimOp mux,
1990 mlir::PatternRewriter &rewriter) const override {
1991 auto width = mux.getType().getBitWidthOrSentinel();
1992 if (width < 0)
1993 return failure();
1994
1995 if (Value v = tryCondTrue(mux.getHigh(), mux.getSel(), rewriter, true, 0)) {
1996 rewriter.modifyOpInPlace(mux, [&] { mux.setOperand(1, v); });
1997 return success();
1998 }
1999
2000 if (Value v = tryCondFalse(mux.getLow(), mux.getSel(), rewriter, true, 0)) {
2001 rewriter.modifyOpInPlace(mux, [&] { mux.setOperand(2, v); });
2002 return success();
2003 }
2004
2005 return failure();
2006 }
2007};
2008} // namespace
2009
2010void MuxPrimOp::getCanonicalizationPatterns(RewritePatternSet &results,
2011 MLIRContext *context) {
2012 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>(
2017 context);
2018}
2019
2020void Mux2CellIntrinsicOp::getCanonicalizationPatterns(
2021 RewritePatternSet &results, MLIRContext *context) {
2022 results.add<patterns::Mux2PadSel>(context);
2023}
2024
2025void Mux4CellIntrinsicOp::getCanonicalizationPatterns(
2026 RewritePatternSet &results, MLIRContext *context) {
2027 results.add<patterns::Mux4PadSel>(context);
2028}
2029
2030OpFoldResult PadPrimOp::fold(FoldAdaptor adaptor) {
2031 auto input = this->getInput();
2032
2033 // pad(x) -> x if the width doesn't change.
2034 if (input.getType() == getType())
2035 return input;
2036
2037 // Need to know the input width.
2038 auto inputType = input.getType().base();
2039 int32_t width = inputType.getWidthOrSentinel();
2040 if (width == -1)
2041 return {};
2042
2043 // Constant fold.
2044 if (auto cst = getConstant(adaptor.getInput())) {
2045 auto destWidth = getType().base().getWidthOrSentinel();
2046 if (destWidth == -1)
2047 return {};
2048
2049 if (inputType.isSigned() && cst->getBitWidth())
2050 return getIntAttr(getType(), cst->sext(destWidth));
2051 return getIntAttr(getType(), cst->zext(destWidth));
2052 }
2053
2054 return {};
2055}
2056
2057OpFoldResult ShlPrimOp::fold(FoldAdaptor adaptor) {
2058 auto input = this->getInput();
2059 IntType inputType = input.getType();
2060 int shiftAmount = getAmount();
2061
2062 // shl(x, 0) -> x
2063 if (shiftAmount == 0)
2064 return input;
2065
2066 // Constant fold.
2067 if (auto cst = getConstant(adaptor.getInput())) {
2068 auto inputWidth = inputType.getWidthOrSentinel();
2069 if (inputWidth != -1) {
2070 auto resultWidth = inputWidth + shiftAmount;
2071 shiftAmount = std::min(shiftAmount, resultWidth);
2072 return getIntAttr(getType(), cst->zext(resultWidth).shl(shiftAmount));
2073 }
2074 }
2075 return {};
2076}
2077
2078OpFoldResult ShrPrimOp::fold(FoldAdaptor adaptor) {
2079 auto input = this->getInput();
2080 IntType inputType = input.getType();
2081 int shiftAmount = getAmount();
2082 auto inputWidth = inputType.getWidthOrSentinel();
2083
2084 // shr(x, 0) -> x
2085 // Once the shr width changes, do this: shiftAmount == 0 &&
2086 // (!inputType.isSigned() || inputWidth > 0)
2087 if (shiftAmount == 0 && inputWidth > 0)
2088 return input;
2089
2090 if (inputWidth == -1)
2091 return {};
2092 if (inputWidth == 0)
2093 return getIntZerosAttr(getType());
2094
2095 // shr(x, cst) where cst is all of x's bits and x is unsigned is 0.
2096 // If x is signed, it is the sign bit.
2097 if (shiftAmount >= inputWidth && inputType.isUnsigned())
2098 return getIntAttr(getType(), APInt(0, 0, false));
2099
2100 // Constant fold.
2101 if (auto cst = getConstant(adaptor.getInput())) {
2102 APInt value;
2103 if (inputType.isSigned())
2104 value = cst->ashr(std::min(shiftAmount, inputWidth - 1));
2105 else
2106 value = cst->lshr(std::min(shiftAmount, inputWidth));
2107 auto resultWidth = std::max(inputWidth - shiftAmount, 1);
2108 return getIntAttr(getType(), value.trunc(resultWidth));
2109 }
2110 return {};
2111}
2112
2113LogicalResult ShrPrimOp::canonicalize(ShrPrimOp op, PatternRewriter &rewriter) {
2114 auto inputWidth = op.getInput().getType().base().getWidthOrSentinel();
2115 if (inputWidth <= 0)
2116 return failure();
2117
2118 // If we know the input width, we can canonicalize this into a BitsPrimOp.
2119 unsigned shiftAmount = op.getAmount();
2120 if (int(shiftAmount) >= inputWidth) {
2121 // shift(x, 32) => 0 when x has 32 bits. This is handled by fold().
2122 if (op.getType().base().isUnsigned())
2123 return failure();
2124
2125 // Shifting a signed value by the full width is actually taking the
2126 // sign bit. If the shift amount is greater than the input width, it
2127 // is equivalent to shifting by the input width.
2128 shiftAmount = inputWidth - 1;
2129 }
2130
2131 replaceWithBits(op, op.getInput(), inputWidth - 1, shiftAmount, rewriter);
2132 return success();
2133}
2134
2135LogicalResult HeadPrimOp::canonicalize(HeadPrimOp op,
2136 PatternRewriter &rewriter) {
2137 auto inputWidth = op.getInput().getType().base().getWidthOrSentinel();
2138 if (inputWidth <= 0)
2139 return failure();
2140
2141 // If we know the input width, we can canonicalize this into a BitsPrimOp.
2142 unsigned keepAmount = op.getAmount();
2143 if (keepAmount)
2144 replaceWithBits(op, op.getInput(), inputWidth - 1, inputWidth - keepAmount,
2145 rewriter);
2146 return success();
2147}
2148
2149OpFoldResult HeadPrimOp::fold(FoldAdaptor adaptor) {
2150 if (hasKnownWidthIntTypes(*this))
2151 if (auto cst = getConstant(adaptor.getInput())) {
2152 int shiftAmount =
2153 getInput().getType().base().getWidthOrSentinel() - getAmount();
2154 return getIntAttr(getType(), cst->lshr(shiftAmount).trunc(getAmount()));
2155 }
2156
2157 return {};
2158}
2159
2160OpFoldResult TailPrimOp::fold(FoldAdaptor adaptor) {
2161 if (hasKnownWidthIntTypes(*this))
2162 if (auto cst = getConstant(adaptor.getInput()))
2163 return getIntAttr(getType(),
2164 cst->trunc(getType().base().getWidthOrSentinel()));
2165 return {};
2166}
2167
2168LogicalResult TailPrimOp::canonicalize(TailPrimOp op,
2169 PatternRewriter &rewriter) {
2170 auto inputWidth = op.getInput().getType().base().getWidthOrSentinel();
2171 if (inputWidth <= 0)
2172 return failure();
2173
2174 // If we know the input width, we can canonicalize this into a BitsPrimOp.
2175 unsigned dropAmount = op.getAmount();
2176 if (dropAmount != unsigned(inputWidth))
2177 replaceWithBits(op, op.getInput(), inputWidth - dropAmount - 1, 0,
2178 rewriter);
2179 return success();
2180}
2181
2182void SubaccessOp::getCanonicalizationPatterns(RewritePatternSet &results,
2183 MLIRContext *context) {
2184 results.add<patterns::SubaccessOfConstant>(context);
2185}
2186
2187OpFoldResult MultibitMuxOp::fold(FoldAdaptor adaptor) {
2188 // If there is only one input, just return it.
2189 if (adaptor.getInputs().size() == 1)
2190 return getOperand(1);
2191
2192 if (auto constIndex = getConstant(adaptor.getIndex())) {
2193 auto index = constIndex->getZExtValue();
2194 if (index < getInputs().size())
2195 return getInputs()[getInputs().size() - 1 - index];
2196 }
2197
2198 return {};
2199}
2200
2201LogicalResult MultibitMuxOp::canonicalize(MultibitMuxOp op,
2202 PatternRewriter &rewriter) {
2203 // If all operands are equal, just canonicalize to it. We can add this
2204 // canonicalization as a folder but it costly to look through all inputs so it
2205 // is added here.
2206 if (llvm::all_of(op.getInputs().drop_front(), [&](auto input) {
2207 return input == op.getInputs().front();
2208 })) {
2209 replaceOpAndCopyName(rewriter, op, op.getInputs().front());
2210 return success();
2211 }
2212
2213 // If the index width is narrower than the size of inputs, drop front
2214 // elements.
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));
2220 });
2221 return success();
2222 }
2223
2224 // If the op is a vector indexing (e.g. `multbit_mux idx, a[n-1], a[n-2], ...,
2225 // a[0]`), we can fold the op into subaccess op `a[idx]`.
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();
2231 })) {
2232 replaceOpWithNewOpAndCopyName<SubaccessOp>(
2233 rewriter, op, lastSubindex.getInput(), op.getIndex());
2234 return success();
2235 }
2236 }
2237
2238 // If the size is 2, canonicalize into a normal mux to introduce more folds.
2239 if (op.getInputs().size() != 2)
2240 return failure();
2241
2242 // TODO: Handle even when `index` doesn't have uint<1>.
2243 auto uintType = op.getIndex().getType();
2244 if (uintType.getBitWidthOrSentinel() != 1)
2245 return failure();
2246
2247 // multibit_mux(index, {lhs, rhs}) -> mux(index, lhs, rhs)
2248 replaceOpWithNewOpAndCopyName<MuxPrimOp>(
2249 rewriter, op, op.getIndex(), op.getInputs()[0], op.getInputs()[1]);
2250 return success();
2251}
2252
2253//===----------------------------------------------------------------------===//
2254// Declarations
2255//===----------------------------------------------------------------------===//
2256
2257/// Scan all the uses of the specified value, checking to see if there is
2258/// exactly one connect that has the value as its destination. This returns the
2259/// operation if found and if all the other users are "reads" from the value.
2260/// Returns null if there are no connects, or multiple connects to the value, or
2261/// if the value is involved in an `AttachOp`, or if the connect isn't matching.
2262///
2263/// Note that this will simply return the connect, which is located *anywhere*
2264/// after the definition of the value. Users of this function are likely
2265/// interested in the source side of the returned connect, the definition of
2266/// which does likely not dominate the original value.
2267MatchingConnectOp firrtl::getSingleConnectUserOf(Value value) {
2268 MatchingConnectOp connect;
2269 for (Operation *user : value.getUsers()) {
2270 // If we see an attach or aggregate sublements, just conservatively fail.
2271 if (isa<AttachOp, SubfieldOp, SubaccessOp, SubindexOp>(user))
2272 return {};
2273
2274 if (auto aConnect = dyn_cast<FConnectLike>(user))
2275 if (aConnect.getDest() == value) {
2276 auto matchingConnect = dyn_cast<MatchingConnectOp>(*aConnect);
2277 // If this is not a matching connect, a second matching connect or in a
2278 // different block, fail.
2279 if (!matchingConnect || (connect && connect != matchingConnect) ||
2280 matchingConnect->getBlock() != value.getParentBlock())
2281 return {};
2282 connect = matchingConnect;
2283 }
2284 }
2285 return connect;
2286}
2287
2288// Forward simple values through wire's and reg's.
2289static LogicalResult canonicalizeSingleSetConnect(MatchingConnectOp op,
2290 PatternRewriter &rewriter) {
2291 // While we can do this for nearly all wires, we currently limit it to simple
2292 // things.
2293 Operation *connectedDecl = op.getDest().getDefiningOp();
2294 if (!connectedDecl)
2295 return failure();
2296
2297 // Only support wire and reg for now.
2298 if (!isa<WireOp>(connectedDecl) && !isa<RegOp>(connectedDecl))
2299 return failure();
2300 if (hasDontTouch(connectedDecl) || !AnnotationSet(connectedDecl).empty() ||
2301 !hasDroppableName(connectedDecl) ||
2302 cast<Forceable>(connectedDecl).isForceable())
2303 return failure();
2304
2305 // Only forward if the types exactly match and there is one connect.
2306 if (getSingleConnectUserOf(op.getDest()) != op)
2307 return failure();
2308
2309 // Only forward if there is more than one use
2310 if (connectedDecl->hasOneUse())
2311 return failure();
2312
2313 // Only do this if the connectee and the declaration are in the same block.
2314 auto *declBlock = connectedDecl->getBlock();
2315 auto *srcValueOp = op.getSrc().getDefiningOp();
2316 if (!srcValueOp) {
2317 // Ports are ok for wires but not registers.
2318 if (!isa<WireOp>(connectedDecl))
2319 return failure();
2320
2321 } else {
2322 // Constants/invalids in the same block are ok to forward, even through
2323 // reg's since the clocking doesn't matter for constants.
2324 auto cnst = dyn_cast<ConstantOp>(srcValueOp);
2325 if (!cnst)
2326 return failure();
2327 if (srcValueOp->getBlock() != declBlock)
2328 return failure();
2329 // A register with a time-zero `initial` value only reaches the constant on
2330 // the first edge; until then it holds `initial`. Forwarding the constant
2331 // would change its value at time zero, so bail unless the two agree.
2332 if (auto reg = dyn_cast<RegOp>(connectedDecl))
2333 if (!preservesInitial(reg.getInitialAttr(), cnst.getValue()))
2334 return failure();
2335 }
2336
2337 // Ok, we know we are doing the transformation.
2338
2339 auto replacement = op.getSrc();
2340 // This will be replaced with the constant source. First, make sure the
2341 // constant dominates all users.
2342 if (srcValueOp && srcValueOp != &declBlock->front())
2343 srcValueOp->moveBefore(&declBlock->front());
2344
2345 // Replace all things *using* the decl with the constant/port, and
2346 // remove the declaration.
2347 replaceOpAndCopyName(rewriter, connectedDecl, replacement);
2348
2349 // Remove the connect
2350 rewriter.eraseOp(op);
2351 return success();
2352}
2353
2354void ConnectOp::getCanonicalizationPatterns(RewritePatternSet &results,
2355 MLIRContext *context) {
2356 results.insert<patterns::ConnectExtension, patterns::ConnectSameType>(
2357 context);
2358}
2359
2360LogicalResult MatchingConnectOp::canonicalize(MatchingConnectOp op,
2361 PatternRewriter &rewriter) {
2362 // TODO: Canonicalize towards explicit extensions and flips here.
2363
2364 // If there is a simple value connected to a foldable decl like a wire or reg,
2365 // see if we can eliminate the decl.
2366 if (succeeded(canonicalizeSingleSetConnect(op, rewriter)))
2367 return success();
2368 return failure();
2369}
2370
2371//===----------------------------------------------------------------------===//
2372// Statements
2373//===----------------------------------------------------------------------===//
2374
2375/// If the specified value has an AttachOp user strictly dominating by
2376/// "dominatingAttach" then return it.
2377static AttachOp getDominatingAttachUser(Value value, AttachOp dominatedAttach) {
2378 for (auto *user : value.getUsers()) {
2379 auto attach = dyn_cast<AttachOp>(user);
2380 if (!attach || attach == dominatedAttach)
2381 continue;
2382 if (attach->isBeforeInBlock(dominatedAttach))
2383 return attach;
2384 }
2385 return {};
2386}
2387
2388LogicalResult AttachOp::canonicalize(AttachOp op, PatternRewriter &rewriter) {
2389 // Single operand attaches are a noop.
2390 if (op.getNumOperands() <= 1) {
2391 rewriter.eraseOp(op);
2392 return success();
2393 }
2394
2395 for (auto operand : op.getOperands()) {
2396 // Check to see if any of our operands has other attaches to it:
2397 // attach x, y
2398 // ...
2399 // attach x, z
2400 // If so, we can merge these into "attach x, y, z".
2401 if (auto attach = getDominatingAttachUser(operand, op)) {
2402 SmallVector<Value> newOperands(op.getOperands());
2403 for (auto newOperand : attach.getOperands())
2404 if (newOperand != operand) // Don't add operand twice.
2405 newOperands.push_back(newOperand);
2406 AttachOp::create(rewriter, op->getLoc(), newOperands);
2407 rewriter.eraseOp(attach);
2408 rewriter.eraseOp(op);
2409 return success();
2410 }
2411
2412 // If this wire is *only* used by an attach then we can just delete
2413 // it.
2414 // TODO: May need to be sensitive to "don't touch" or other
2415 // annotations.
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) // Don't the add wire.
2422 newOperands.push_back(newOperand);
2423
2424 AttachOp::create(rewriter, op->getLoc(), newOperands);
2425 rewriter.eraseOp(op);
2426 rewriter.eraseOp(wire);
2427 return success();
2428 }
2429 }
2430 }
2431 return failure();
2432}
2433
2434/// Replaces the given op with the contents of the given single-block region.
2435static void replaceOpWithRegion(PatternRewriter &rewriter, Operation *op,
2436 Region &region) {
2437 assert(llvm::hasSingleElement(region) && "expected single-region block");
2438 rewriter.inlineBlockBefore(&region.front(), op, {});
2439}
2440
2441LogicalResult WhenOp::canonicalize(WhenOp op, PatternRewriter &rewriter) {
2442 if (auto constant = op.getCondition().getDefiningOp<firrtl::ConstantOp>()) {
2443 if (constant.getValue().isAllOnes())
2444 replaceOpWithRegion(rewriter, op, op.getThenRegion());
2445 else if (op.hasElseRegion() && !op.getElseRegion().empty())
2446 replaceOpWithRegion(rewriter, op, op.getElseRegion());
2447
2448 rewriter.eraseOp(op);
2449
2450 return success();
2451 }
2452
2453 // Erase empty if-else block.
2454 if (!op.getThenBlock().empty() && op.hasElseRegion() &&
2455 op.getElseBlock().empty()) {
2456 rewriter.eraseBlock(&op.getElseBlock());
2457 return success();
2458 }
2459
2460 // Erase empty whens.
2461
2462 // If there is stuff in the then block, leave this operation alone.
2463 if (!op.getThenBlock().empty())
2464 return failure();
2465
2466 // If not and there is no else, then this operation is just useless.
2467 if (!op.hasElseRegion() || op.getElseBlock().empty()) {
2468 rewriter.eraseOp(op);
2469 return success();
2470 }
2471 return failure();
2472}
2473
2474namespace {
2475// Remove private nodes. If they have an interesting names, move the name to
2476// the source expression.
2477struct FoldNodeName : public mlir::OpRewritePattern<NodeOp> {
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() ||
2483 !AnnotationSet(node).empty() || node.isForceable())
2484 return failure();
2485 auto *newOp = node.getInput().getDefiningOp();
2486 if (newOp)
2487 updateName(rewriter, newOp, name);
2488 rewriter.replaceOp(node, node.getInput());
2489 return success();
2490 }
2491};
2492
2493// Bypass nodes.
2494struct NodeBypass : public mlir::OpRewritePattern<NodeOp> {
2495 using OpRewritePattern::OpRewritePattern;
2496 LogicalResult matchAndRewrite(NodeOp node,
2497 PatternRewriter &rewriter) const override {
2498 if (node.getInnerSym() || !AnnotationSet(node).empty() ||
2499 node.use_empty() || node.isForceable())
2500 return failure();
2501 rewriter.replaceAllUsesWith(node.getResult(), node.getInput());
2502 return success();
2503 }
2504};
2505
2506} // namespace
2507
2508template <typename OpTy>
2509static LogicalResult demoteForceableIfUnused(OpTy op,
2510 PatternRewriter &rewriter) {
2511 if (!op.isForceable() || !op.getDataRef().use_empty())
2512 return failure();
2513
2514 firrtl::detail::replaceWithNewForceability(op, false, &rewriter);
2515 return success();
2516}
2517
2518// Interesting names and symbols and don't touch force nodes to stick around.
2519LogicalResult NodeOp::fold(FoldAdaptor adaptor,
2520 SmallVectorImpl<OpFoldResult> &results) {
2521 if (!hasDroppableName())
2522 return failure();
2523 if (hasDontTouch(getResult())) // handles inner symbols
2524 return failure();
2525 if (getAnnotationsAttr() && !AnnotationSet(getAnnotationsAttr()).empty())
2526 return failure();
2527 if (isForceable())
2528 return failure();
2529 if (!adaptor.getInput())
2530 return failure();
2531
2532 results.push_back(adaptor.getInput());
2533 return success();
2534}
2535
2536void NodeOp::getCanonicalizationPatterns(RewritePatternSet &results,
2537 MLIRContext *context) {
2538 results.insert<FoldNodeName>(context);
2539 results.add(demoteForceableIfUnused<NodeOp>);
2540}
2541
2542namespace {
2543// For a lhs, find all the writers of fields of the aggregate type. If there
2544// is one writer for each field, merge the writes
2545struct AggOneShot : public mlir::RewritePattern {
2546 AggOneShot(StringRef name, uint32_t weight, MLIRContext *context)
2547 : RewritePattern(name, 0, context) {}
2548
2549 SmallVector<Value> getCompleteWrite(Operation *lhs) const {
2550 auto lhsTy = lhs->getResult(0).getType();
2551 if (!type_isa<BundleType, FVectorType>(lhsTy))
2552 return {};
2553
2554 DenseMap<uint32_t, Value> fields;
2555 for (Operation *user : lhs->getResult(0).getUsers()) {
2556 if (user->getParentOp() != lhs->getParentOp())
2557 return {};
2558 if (auto aConnect = dyn_cast<MatchingConnectOp>(user)) {
2559 if (aConnect.getDest() == lhs->getResult(0))
2560 return {};
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())
2566 return {};
2567 if (fields.count(subField.getFieldIndex())) // duplicate write
2568 return {};
2569 fields[subField.getFieldIndex()] = aConnect.getSrc();
2570 }
2571 continue;
2572 }
2573 return {};
2574 }
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())
2580 return {};
2581 if (fields.count(subIndex.getIndex())) // duplicate write
2582 return {};
2583 fields[subIndex.getIndex()] = aConnect.getSrc();
2584 }
2585 continue;
2586 }
2587 return {};
2588 }
2589 } else {
2590 return {};
2591 }
2592 }
2593
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))
2600 return {};
2601 values.push_back(fields[i]);
2602 }
2603 return values;
2604 }
2605
2606 LogicalResult matchAndRewrite(Operation *op,
2607 PatternRewriter &rewriter) const override {
2608 auto values = getCompleteWrite(op);
2609 if (values.empty())
2610 return failure();
2611 rewriter.setInsertionPointToEnd(op->getBlock());
2612 auto dest = op->getResult(0);
2613 auto destType = dest.getType();
2614
2615 // If not passive, cannot matchingconnect.
2616 if (!type_cast<FIRRTLBaseType>(destType).isPassive())
2617 return failure();
2618
2619 Value newVal = type_isa<BundleType>(destType)
2620 ? rewriter.createOrFold<BundleCreateOp>(op->getLoc(),
2621 destType, values)
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);
2638 }
2639 }
2640 return success();
2641 }
2642};
2643
2644struct WireAggOneShot : public AggOneShot {
2645 WireAggOneShot(MLIRContext *context)
2646 : AggOneShot(WireOp::getOperationName(), 0, context) {}
2647};
2648struct SubindexAggOneShot : public AggOneShot {
2649 SubindexAggOneShot(MLIRContext *context)
2650 : AggOneShot(SubindexOp::getOperationName(), 0, context) {}
2651};
2652struct SubfieldAggOneShot : public AggOneShot {
2653 SubfieldAggOneShot(MLIRContext *context)
2654 : AggOneShot(SubfieldOp::getOperationName(), 0, context) {}
2655};
2656} // namespace
2657
2658void WireOp::getCanonicalizationPatterns(RewritePatternSet &results,
2659 MLIRContext *context) {
2660 results.insert<WireAggOneShot>(context);
2661 results.add(demoteForceableIfUnused<WireOp>);
2662}
2663
2664void SubindexOp::getCanonicalizationPatterns(RewritePatternSet &results,
2665 MLIRContext *context) {
2666 results.insert<SubindexAggOneShot>(context);
2667}
2668
2669OpFoldResult SubindexOp::fold(FoldAdaptor adaptor) {
2670 auto attr = dyn_cast_or_null<ArrayAttr>(adaptor.getInput());
2671 if (!attr)
2672 return {};
2673 return attr[getIndex()];
2674}
2675
2676OpFoldResult SubfieldOp::fold(FoldAdaptor adaptor) {
2677 auto attr = dyn_cast_or_null<ArrayAttr>(adaptor.getInput());
2678 if (!attr)
2679 return {};
2680 auto index = getFieldIndex();
2681 return attr[index];
2682}
2683
2684void SubfieldOp::getCanonicalizationPatterns(RewritePatternSet &results,
2685 MLIRContext *context) {
2686 results.insert<SubfieldAggOneShot>(context);
2687}
2688
2689static Attribute collectFields(MLIRContext *context,
2690 ArrayRef<Attribute> operands) {
2691 for (auto operand : operands)
2692 if (!operand)
2693 return {};
2694 return ArrayAttr::get(context, operands);
2695}
2696
2697OpFoldResult BundleCreateOp::fold(FoldAdaptor adaptor) {
2698 // bundle_create(%foo["a"], %foo["b"]) -> %foo when the type of %foo is
2699 // bundle<a:..., b:...>.
2700 if (getNumOperands() > 0)
2701 if (SubfieldOp first = getOperand(0).getDefiningOp<SubfieldOp>())
2702 if (first.getFieldIndex() == 0 &&
2703 first.getInput().getType() == getType() &&
2704 llvm::all_of(
2705 llvm::drop_begin(llvm::enumerate(getOperands())), [&](auto elem) {
2706 auto subindex =
2707 elem.value().template getDefiningOp<SubfieldOp>();
2708 return subindex && subindex.getInput() == first.getInput() &&
2709 subindex.getFieldIndex() == elem.index();
2710 }))
2711 return first.getInput();
2712
2713 return collectFields(getContext(), adaptor.getOperands());
2714}
2715
2716OpFoldResult VectorCreateOp::fold(FoldAdaptor adaptor) {
2717 // vector_create(%foo[0], %foo[1]) -> %foo when the type of %foo is
2718 // vector<..., 2>.
2719 if (getNumOperands() > 0)
2720 if (SubindexOp first = getOperand(0).getDefiningOp<SubindexOp>())
2721 if (first.getIndex() == 0 && first.getInput().getType() == getType() &&
2722 llvm::all_of(
2723 llvm::drop_begin(llvm::enumerate(getOperands())), [&](auto elem) {
2724 auto subindex =
2725 elem.value().template getDefiningOp<SubindexOp>();
2726 return subindex && subindex.getInput() == first.getInput() &&
2727 subindex.getIndex() == elem.index();
2728 }))
2729 return first.getInput();
2730
2731 return collectFields(getContext(), adaptor.getOperands());
2732}
2733
2734OpFoldResult UninferredResetCastOp::fold(FoldAdaptor adaptor) {
2735 if (getOperand().getType() == getType())
2736 return getOperand();
2737 return {};
2738}
2739
2740namespace {
2741// A register with constant reset and all connection to either itself or the
2742// same constant, must be replaced by the constant.
2743struct FoldResetMux : public mlir::OpRewritePattern<RegResetOp> {
2744 using OpRewritePattern::OpRewritePattern;
2745 LogicalResult matchAndRewrite(RegResetOp reg,
2746 PatternRewriter &rewriter) const override {
2747 auto reset =
2748 dyn_cast_or_null<ConstantOp>(reg.getResetValue().getDefiningOp());
2749 if (!reset || hasDontTouch(reg.getOperation()) ||
2750 !AnnotationSet(reg).empty() || reg.isForceable())
2751 return failure();
2752 // Find the one true connect, or bail
2753 auto con = getSingleConnectUserOf(reg.getResult());
2754 if (!con)
2755 return failure();
2756
2757 auto mux = dyn_cast_or_null<MuxPrimOp>(con.getSrc().getDefiningOp());
2758 if (!mux)
2759 return failure();
2760 auto *high = mux.getHigh().getDefiningOp();
2761 auto *low = mux.getLow().getDefiningOp();
2762 auto constOp = dyn_cast_or_null<ConstantOp>(high);
2763
2764 if (constOp && low != reg)
2765 return failure();
2766 if (dyn_cast_or_null<ConstantOp>(low) && high == reg)
2767 constOp = dyn_cast<ConstantOp>(low);
2768
2769 if (!constOp || constOp.getType() != reset.getType() ||
2770 constOp.getValue() != reset.getValue())
2771 return failure();
2772
2773 // Check all types should be typed by now
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)
2778 return failure();
2779
2780 // Ok, we know we are doing the transformation.
2781
2782 // Bail if replacing the register with the constant would change its
2783 // time-zero simulation value.
2784 if (!preservesInitial(reg.getInitialAttr(), constOp.getValue()))
2785 return failure();
2786
2787 // Make sure the constant dominates all users.
2788 if (constOp != &con->getBlock()->front())
2789 constOp->moveBefore(&con->getBlock()->front());
2790
2791 // Replace the register with the constant.
2792 replaceOpAndCopyName(rewriter, reg, constOp.getResult());
2793 // Remove the connect.
2794 rewriter.eraseOp(con);
2795 return success();
2796 }
2797};
2798} // namespace
2799
2800static bool isDefinedByOneConstantOp(Value v) {
2801 if (auto c = v.getDefiningOp<ConstantOp>())
2802 return c.getValue().isOne();
2803 if (auto sc = v.getDefiningOp<SpecialConstantOp>())
2804 return sc.getValue();
2805 return false;
2806}
2807
2808static LogicalResult
2809canonicalizeRegResetWithOneReset(RegResetOp reg, PatternRewriter &rewriter) {
2810 if (!isDefinedByOneConstantOp(reg.getResetSignal()))
2811 return failure();
2812
2813 auto resetValue = reg.getResetValue();
2814 if (reg.getType(0) != resetValue.getType())
2815 return failure();
2816
2817 // Bail if replacing the register with the constant would change its
2818 // time-zero simulation value.
2819 if (auto constOp = dyn_cast_or_null<ConstantOp>(resetValue.getDefiningOp())) {
2820 if (!preservesInitial(reg.getInitialAttr(), constOp.getValue()))
2821 return failure();
2822 } else if (!preservesInitial(reg.getInitialAttr())) {
2823 return failure();
2824 }
2825
2826 // Ignore 'passthrough'.
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());
2831 return success();
2832}
2833
2834void RegResetOp::getCanonicalizationPatterns(RewritePatternSet &results,
2835 MLIRContext *context) {
2836 results.add<patterns::RegResetWithZeroReset, FoldResetMux>(context);
2838 results.add(demoteForceableIfUnused<RegResetOp>);
2839}
2840
2841// Returns the value connected to a port, if there is only one.
2842static Value getPortFieldValue(Value port, StringRef name) {
2843 auto portTy = type_cast<BundleType>(port.getType());
2844 auto fieldIndex = portTy.getElementIndex(name);
2845 assert(fieldIndex && "missing field on memory port");
2846
2847 Value value = {};
2848 for (auto *op : port.getUsers()) {
2849 auto portAccess = cast<SubfieldOp>(op);
2850 if (fieldIndex != portAccess.getFieldIndex())
2851 continue;
2852 auto conn = getSingleConnectUserOf(portAccess);
2853 if (!conn || value)
2854 return {};
2855 value = conn.getSrc();
2856 }
2857 return value;
2858}
2859
2860// Returns true if the enable field of a port is set to false.
2861static bool isPortDisabled(Value port) {
2862 auto value = getPortFieldValue(port, "en");
2863 if (!value)
2864 return false;
2865 auto portConst = value.getDefiningOp<ConstantOp>();
2866 if (!portConst)
2867 return false;
2868 return portConst.getValue().isZero();
2869}
2870
2871// Returns true if the data output is unused.
2872static bool isPortUnused(Value port, StringRef data) {
2873 auto portTy = type_cast<BundleType>(port.getType());
2874 auto fieldIndex = portTy.getElementIndex(data);
2875 assert(fieldIndex && "missing enable flag on memory port");
2876
2877 for (auto *op : port.getUsers()) {
2878 auto portAccess = cast<SubfieldOp>(op);
2879 if (fieldIndex != portAccess.getFieldIndex())
2880 continue;
2881 if (!portAccess.use_empty())
2882 return false;
2883 }
2884
2885 return true;
2886}
2887
2888// Returns the value connected to a port, if there is only one.
2889static void replacePortField(PatternRewriter &rewriter, Value port,
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");
2894
2895 for (auto *op : llvm::make_early_inc_range(port.getUsers())) {
2896 auto portAccess = cast<SubfieldOp>(op);
2897 if (fieldIndex != portAccess.getFieldIndex())
2898 continue;
2899 rewriter.replaceAllUsesWith(portAccess, value);
2900 rewriter.eraseOp(portAccess);
2901 }
2902}
2903
2904// Remove accesses to a port which is used.
2905static void erasePort(PatternRewriter &rewriter, Value port) {
2906 // Helper to create a dummy 0 clock for the dummy registers.
2907 Value clock;
2908 auto getClock = [&] {
2909 if (!clock)
2910 clock = SpecialConstantOp::create(rewriter, port.getLoc(),
2911 ClockType::get(rewriter.getContext()),
2912 false);
2913 return clock;
2914 };
2915
2916 // Find the clock field of the port and determine whether the port is
2917 // accessed only through its subfields or as a whole wire. If the port
2918 // is used in its entirety, replace it with a wire. Otherwise,
2919 // eliminate individual subfields and replace with reasonable defaults.
2920 for (auto *op : port.getUsers()) {
2921 auto subfield = dyn_cast<SubfieldOp>(op);
2922 if (!subfield) {
2923 auto ty = port.getType();
2924 auto reg = RegOp::create(rewriter, port.getLoc(), ty, getClock());
2925 rewriter.replaceAllUsesWith(port, reg.getResult());
2926 return;
2927 }
2928 }
2929
2930 // Remove all connects to field accesses as they are no longer relevant.
2931 // If field values are used anywhere, which should happen solely for read
2932 // ports, a dummy register is introduced which replicates the behaviour of
2933 // memory that is never written, but might be read.
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);
2940 continue;
2941 }
2942 }
2943 if (access.use_empty()) {
2944 rewriter.eraseOp(access);
2945 continue;
2946 }
2947
2948 // Replace read values with a register that is never written, handing off
2949 // the canonicalization of such a register to another canonicalizer.
2950 auto ty = access.getType();
2951 auto reg = RegOp::create(rewriter, access.getLoc(), ty, getClock());
2952 rewriter.replaceOp(access, reg.getResult());
2953 }
2954 assert(port.use_empty() && "port should have no remaining uses");
2955}
2956
2957namespace {
2958// If memory has known, but zero width, eliminate it.
2959struct FoldZeroWidthMemory : public mlir::OpRewritePattern<MemOp> {
2960 using OpRewritePattern::OpRewritePattern;
2961 LogicalResult matchAndRewrite(MemOp mem,
2962 PatternRewriter &rewriter) const override {
2963 if (hasDontTouch(mem))
2964 return failure();
2965
2966 if (!firrtl::type_isa<IntType>(mem.getDataType()) ||
2967 mem.getDataType().getBitWidthOrSentinel() != 0)
2968 return failure();
2969
2970 // Make sure are users are safe to replace
2971 for (auto port : mem.getResults())
2972 for (auto *user : port.getUsers())
2973 if (!isa<SubfieldOp>(user))
2974 return failure();
2975
2976 // Annoyingly, there isn't a good replacement for the port as a whole,
2977 // since they have an outer flip type.
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())
2984 .getResult();
2985 if (fieldName.ends_with("data")) {
2986 // Make sure to write data ports.
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);
2991 }
2992 }
2993 }
2994 rewriter.eraseOp(mem);
2995 return success();
2996 }
2997};
2998
2999// If memory has no write ports and no file initialization, eliminate it.
3000struct FoldReadOrWriteOnlyMemory : public mlir::OpRewritePattern<MemOp> {
3001 using OpRewritePattern::OpRewritePattern;
3002 LogicalResult matchAndRewrite(MemOp mem,
3003 PatternRewriter &rewriter) const override {
3004 if (hasDontTouch(mem))
3005 return failure();
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:
3010 isRead = true;
3011 if (isWritten)
3012 return failure();
3013 continue;
3014 case MemOp::PortKind::Write:
3015 isWritten = true;
3016 if (isRead)
3017 return failure();
3018 continue;
3019 case MemOp::PortKind::Debug:
3020 case MemOp::PortKind::ReadWrite:
3021 return failure();
3022 }
3023 llvm_unreachable("unknown port kind");
3024 }
3025 assert((!isWritten || !isRead) && "memory is in use");
3026
3027 // If the memory is read only, but has a file initialization, then we can't
3028 // remove it. A write only memory with file initialization is okay to
3029 // remove.
3030 if (isRead && mem.getInit())
3031 return failure();
3032
3033 for (auto port : mem.getResults())
3034 erasePort(rewriter, port);
3035
3036 rewriter.eraseOp(mem);
3037 return success();
3038 }
3039};
3040
3041// Eliminate the dead ports of memories.
3042struct FoldUnusedPorts : public mlir::OpRewritePattern<MemOp> {
3043 using OpRewritePattern::OpRewritePattern;
3044 LogicalResult matchAndRewrite(MemOp mem,
3045 PatternRewriter &rewriter) const override {
3046 if (hasDontTouch(mem))
3047 return failure();
3048 // Identify the dead and changed ports.
3049 llvm::SmallBitVector deadPorts(mem.getNumResults());
3050 for (auto [i, port] : llvm::enumerate(mem.getResults())) {
3051 // Do not simplify annotated ports.
3052 if (!mem.getPortAnnotation(i).empty())
3053 continue;
3054
3055 // Skip debug ports.
3056 auto kind = mem.getPortKind(i);
3057 if (kind == MemOp::PortKind::Debug)
3058 continue;
3059
3060 // If a port is disabled, always eliminate it.
3061 if (isPortDisabled(port)) {
3062 deadPorts.set(i);
3063 continue;
3064 }
3065 // Eliminate read ports whose outputs are not used.
3066 if (kind == MemOp::PortKind::Read && isPortUnused(port, "data")) {
3067 deadPorts.set(i);
3068 continue;
3069 }
3070 }
3071 if (deadPorts.none())
3072 return failure();
3073
3074 // Rebuild the new memory with the altered ports.
3075 SmallVector<Type> resultTypes;
3076 SmallVector<StringRef> portNames;
3077 SmallVector<Attribute> portAnnotations;
3078 for (auto [i, port] : llvm::enumerate(mem.getResults())) {
3079 if (deadPorts[i])
3080 continue;
3081 resultTypes.push_back(port.getType());
3082 portNames.push_back(mem.getPortName(i));
3083 portAnnotations.push_back(mem.getPortAnnotation(i));
3084 }
3085
3086 MemOp newOp;
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());
3094
3095 // Replace the dead ports with dummy wires.
3096 unsigned nextPort = 0;
3097 for (auto [i, port] : llvm::enumerate(mem.getResults())) {
3098 if (deadPorts[i])
3099 erasePort(rewriter, port);
3100 else
3101 rewriter.replaceAllUsesWith(port, newOp.getResult(nextPort++));
3102 }
3103
3104 rewriter.eraseOp(mem);
3105 return success();
3106 }
3107};
3108
3109// Rewrite write-only read-write ports to write ports.
3110struct FoldReadWritePorts : public mlir::OpRewritePattern<MemOp> {
3111 using OpRewritePattern::OpRewritePattern;
3112 LogicalResult matchAndRewrite(MemOp mem,
3113 PatternRewriter &rewriter) const override {
3114 if (hasDontTouch(mem))
3115 return failure();
3116
3117 // Identify read-write ports whose read end is unused.
3118 llvm::SmallBitVector deadReads(mem.getNumResults());
3119 for (auto [i, port] : llvm::enumerate(mem.getResults())) {
3120 if (mem.getPortKind(i) != MemOp::PortKind::ReadWrite)
3121 continue;
3122 if (!mem.getPortAnnotation(i).empty())
3123 continue;
3124 if (isPortUnused(port, "rdata")) {
3125 deadReads.set(i);
3126 continue;
3127 }
3128 }
3129 if (deadReads.none())
3130 return failure();
3131
3132 SmallVector<Type> resultTypes;
3133 SmallVector<StringRef> portNames;
3134 SmallVector<Attribute> portAnnotations;
3135 for (auto [i, port] : llvm::enumerate(mem.getResults())) {
3136 if (deadReads[i])
3137 resultTypes.push_back(
3138 MemOp::getTypeForPort(mem.getDepth(), mem.getDataType(),
3139 MemOp::PortKind::Write, mem.getMaskBits()));
3140 else
3141 resultTypes.push_back(port.getType());
3142
3143 portNames.push_back(mem.getPortName(i));
3144 portAnnotations.push_back(mem.getPortAnnotation(i));
3145 }
3146
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());
3153
3154 for (unsigned i = 0, n = mem.getNumResults(); i < n; ++i) {
3155 auto result = mem.getResult(i);
3156 auto newResult = newOp.getResult(i);
3157 if (deadReads[i]) {
3158 auto resultPortTy = type_cast<BundleType>(result.getType());
3159
3160 // Rewrite accesses to the old port field to accesses to a
3161 // corresponding field of the new port.
3162 auto replace = [&](StringRef toName, StringRef fromName) {
3163 auto fromFieldIndex = resultPortTy.getElementIndex(fromName);
3164 assert(fromFieldIndex && "missing enable flag on memory port");
3165
3166 auto toField = SubfieldOp::create(rewriter, newResult.getLoc(),
3167 newResult, toName);
3168 for (auto *op : llvm::make_early_inc_range(result.getUsers())) {
3169 auto fromField = cast<SubfieldOp>(op);
3170 if (fromFieldIndex != fromField.getFieldIndex())
3171 continue;
3172 rewriter.replaceOp(fromField, toField.getResult());
3173 }
3174 };
3175
3176 replace("addr", "addr");
3177 replace("en", "en");
3178 replace("clk", "clk");
3179 replace("data", "wdata");
3180 replace("mask", "wmask");
3181
3182 // Remove the wmode field, replacing it with dummy wires.
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())
3187 continue;
3188 rewriter.replaceOpWithNewOp<WireOp>(wmodeField, wmodeField.getType());
3189 }
3190 } else {
3191 rewriter.replaceAllUsesWith(result, newResult);
3192 }
3193 }
3194 rewriter.eraseOp(mem);
3195 return success();
3196 }
3197};
3198
3199// Eliminate the dead ports of memories.
3200struct FoldUnusedBits : public mlir::OpRewritePattern<MemOp> {
3201 using OpRewritePattern::OpRewritePattern;
3202
3203 LogicalResult matchAndRewrite(MemOp mem,
3204 PatternRewriter &rewriter) const override {
3205 if (hasDontTouch(mem))
3206 return failure();
3207
3208 // Only apply the transformation if the memory is not sequential.
3209 const auto &summary = mem.getSummary();
3210 if (summary.isMasked || summary.isSeqMem())
3211 return failure();
3212
3213 auto type = type_dyn_cast<IntType>(mem.getDataType());
3214 if (!type)
3215 return failure();
3216 auto width = type.getBitWidthOrSentinel();
3217 if (width <= 0)
3218 return failure();
3219
3220 llvm::SmallBitVector usedBits(width);
3221 DenseMap<unsigned, unsigned> mapping;
3222
3223 // Find which bits are used out of the users of a read port. This detects
3224 // ports whose data/rdata field is used only through bit select ops. The
3225 // bit selects are then used to build a bit-mask. The ops are collected.
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");
3231
3232 for (auto *op : port.getUsers()) {
3233 auto portAccess = cast<SubfieldOp>(op);
3234 if (fieldIndex != portAccess.getFieldIndex())
3235 continue;
3236
3237 for (auto *user : op->getUsers()) {
3238 auto bits = dyn_cast<BitsPrimOp>(user);
3239 if (!bits)
3240 return failure();
3241
3242 usedBits.set(bits.getLo(), bits.getHi() + 1);
3243 if (usedBits.all())
3244 return failure();
3245
3246 mapping[bits.getLo()] = 0;
3247 readOps.push_back(bits);
3248 }
3249 }
3250
3251 return success();
3252 };
3253
3254 // Finds the users of write ports. This expects all the data/wdata fields
3255 // of the ports to be used solely as the destination of matching connects.
3256 // If a memory has ports with other uses, it is excluded from optimisation.
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");
3262
3263 for (auto *op : port.getUsers()) {
3264 auto portAccess = cast<SubfieldOp>(op);
3265 if (fieldIndex != portAccess.getFieldIndex())
3266 continue;
3267
3268 auto conn = getSingleConnectUserOf(portAccess);
3269 if (!conn)
3270 return failure();
3271
3272 writeOps.push_back(conn);
3273 }
3274 return success();
3275 };
3276
3277 // Traverse all ports and find the read and used data fields.
3278 for (auto [i, port] : llvm::enumerate(mem.getResults())) {
3279 // Do not simplify annotated ports.
3280 if (!mem.getPortAnnotation(i).empty())
3281 return failure();
3282
3283 switch (mem.getPortKind(i)) {
3284 case MemOp::PortKind::Debug:
3285 // Skip debug ports.
3286 return failure();
3287 case MemOp::PortKind::Write:
3288 if (failed(findWriteUsers(port, "data")))
3289 return failure();
3290 continue;
3291 case MemOp::PortKind::Read:
3292 if (failed(findReadUsers(port, "data")))
3293 return failure();
3294 continue;
3295 case MemOp::PortKind::ReadWrite:
3296 if (failed(findWriteUsers(port, "wdata")))
3297 return failure();
3298 if (failed(findReadUsers(port, "rdata")))
3299 return failure();
3300 continue;
3301 }
3302 llvm_unreachable("unknown port kind");
3303 }
3304
3305 // Unused memories are handled in a different canonicalizer.
3306 if (usedBits.none())
3307 return failure();
3308
3309 // Build a mapping of existing indices to compacted ones.
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);
3314 if (e < 0)
3315 e = width;
3316 for (int idx = i; idx < e; ++idx, ++newWidth) {
3317 if (auto it = mapping.find(idx); it != mapping.end()) {
3318 it->second = newWidth;
3319 }
3320 }
3321 ranges.emplace_back(i, e - 1);
3322 i = e != width ? usedBits.find_next(e) : e;
3323 }
3324
3325 // Create the new op with the new port types.
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)));
3331 }
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());
3337
3338 // Rewrite bundle users to the new data type.
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");
3343
3344 rewriter.setInsertionPointAfter(newMem);
3345 auto newPortAccess =
3346 SubfieldOp::create(rewriter, port.getLoc(), port, field);
3347
3348 for (auto *op : llvm::make_early_inc_range(port.getUsers())) {
3349 auto portAccess = cast<SubfieldOp>(op);
3350 if (op == newPortAccess || fieldIndex != portAccess.getFieldIndex())
3351 continue;
3352 rewriter.replaceOp(portAccess, newPortAccess.getResult());
3353 }
3354 };
3355
3356 // Rewrite the field accesses.
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");
3363 continue;
3364 case MemOp::PortKind::Read:
3365 rewriteSubfield(port, "data");
3366 continue;
3367 case MemOp::PortKind::ReadWrite:
3368 rewriteSubfield(port, "rdata");
3369 rewriteSubfield(port, "wdata");
3370 continue;
3371 }
3372 llvm_unreachable("unknown port kind");
3373 }
3374
3375 // Rewrite the reads to the new ranges, compacting them.
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");
3380 // Create a new bit selection from the compressed memory. The new op may
3381 // be folded if we are selecting the entire compressed memory.
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);
3387 }
3388
3389 // Rewrite the writes into a concatenation of slices.
3390 for (auto writeOp : writeOps) {
3391 Value source = writeOp.getSrc();
3392 rewriter.setInsertionPoint(writeOp);
3393
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);
3399 }
3400
3401 Value catOfSlices =
3402 rewriter.createOrFold<CatPrimOp>(writeOp.getLoc(), slices);
3403
3404 // If the original memory held a signed integer, then the compressed
3405 // memory will be signed too. Since the catOfSlices is always unsigned,
3406 // cast the data to a signed integer if needed before connecting back to
3407 // the memory.
3408 if (type.isSigned())
3409 catOfSlices =
3410 rewriter.createOrFold<AsSIntPrimOp>(writeOp.getLoc(), catOfSlices);
3411
3412 rewriter.replaceOpWithNewOp<MatchingConnectOp>(writeOp, writeOp.getDest(),
3413 catOfSlices);
3414 }
3415
3416 return success();
3417 }
3418};
3419
3420// Rewrite single-address memories to a firrtl register.
3421struct FoldRegMems : public mlir::OpRewritePattern<MemOp> {
3422 using OpRewritePattern::OpRewritePattern;
3423 LogicalResult matchAndRewrite(MemOp mem,
3424 PatternRewriter &rewriter) const override {
3425 const FirMemory &info = mem.getSummary();
3426 if (hasDontTouch(mem) || info.depth != 1)
3427 return failure();
3428
3429 auto ty = mem.getDataType();
3430 auto loc = mem.getLoc();
3431 auto *block = mem->getBlock();
3432
3433 // Find the clock of the register-to-be, all write ports should share it.
3434 Value clock;
3435 SmallPtrSet<Operation *, 8> connects;
3436 SmallVector<SubfieldOp> portAccesses;
3437 for (auto [i, port] : llvm::enumerate(mem.getResults())) {
3438 if (!mem.getPortAnnotation(i).empty())
3439 continue;
3440
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");
3446
3447 for (auto *op : port.getUsers()) {
3448 auto portAccess = cast<SubfieldOp>(op);
3449 if (fieldIndex != portAccess.getFieldIndex())
3450 continue;
3451 portAccesses.push_back(portAccess);
3452 for (auto *user : portAccess->getUsers()) {
3453 auto conn = dyn_cast<FConnectLike>(user);
3454 if (!conn)
3455 return failure();
3456 connects.insert(conn);
3457 }
3458 }
3459 }
3460 return success();
3461 };
3462
3463 switch (mem.getPortKind(i)) {
3464 case MemOp::PortKind::Debug:
3465 return failure();
3466 case MemOp::PortKind::Read:
3467 if (failed(collect({"clk", "en", "addr"})))
3468 return failure();
3469 continue;
3470 case MemOp::PortKind::Write:
3471 if (failed(collect({"clk", "en", "addr", "data", "mask"})))
3472 return failure();
3473 break;
3474 case MemOp::PortKind::ReadWrite:
3475 if (failed(collect({"clk", "en", "addr", "wmode", "wdata", "wmask"})))
3476 return failure();
3477 break;
3478 }
3479
3480 Value portClock = getPortFieldValue(port, "clk");
3481 if (!portClock || (clock && portClock != clock))
3482 return failure();
3483 clock = portClock;
3484 }
3485 // Create a new wire where the memory used to be. This wire will dominate
3486 // all readers of the memory. Reads should be made through this wire.
3487 rewriter.setInsertionPointAfter(mem);
3488 auto memWire = WireOp::create(rewriter, loc, ty).getResult();
3489
3490 // The memory is replaced by a register, which we place at the end of the
3491 // block, so that any value driven to the original memory will dominate the
3492 // new register (including the clock). All other ops will be placed
3493 // after the register.
3494 rewriter.setInsertionPointToEnd(block);
3495 auto memReg =
3496 RegOp::create(rewriter, loc, ty, clock, mem.getName()).getResult();
3497
3498 // Connect the output of the register to the wire.
3499 MatchingConnectOp::create(rewriter, loc, memWire, memReg);
3500
3501 // Helper to insert a given number of pipeline stages through registers.
3502 // The pipelines are placed at the end of the block.
3503 auto pipeline = [&](Value value, Value clock, const Twine &name,
3504 unsigned latency) {
3505 for (unsigned i = 0; i < latency; ++i) {
3506 std::string regName;
3507 {
3508 llvm::raw_string_ostream os(regName);
3509 os << mem.getName() << "_" << name << "_" << i;
3510 }
3511 auto reg = RegOp::create(rewriter, mem.getLoc(), value.getType(), clock,
3512 rewriter.getStringAttr(regName))
3513 .getResult();
3514 MatchingConnectOp::create(rewriter, value.getLoc(), reg, value);
3515 value = reg;
3516 }
3517 return value;
3518 };
3519
3520 const unsigned writeStages = info.writeLatency - 1;
3521
3522 // Traverse each port. Replace reads with the pipelined register, discarding
3523 // the enable flag and reading unconditionally. Pipeline the mask, enable
3524 // and data bits of all write ports to be arbitrated and wired to the reg.
3525 SmallVector<std::tuple<Value, Value, Value>> writes;
3526 for (auto [i, port] : llvm::enumerate(mem.getResults())) {
3527 Value portClock = getPortFieldValue(port, "clk");
3528 StringRef name = mem.getPortName(i);
3529
3530 auto portPipeline = [&, port = port](StringRef field, unsigned stages) {
3531 Value value = getPortFieldValue(port, field);
3532 assert(value);
3533 return pipeline(value, portClock, name + "_" + field, stages);
3534 };
3535
3536 switch (mem.getPortKind(i)) {
3537 case MemOp::PortKind::Debug:
3538 llvm_unreachable("unknown port kind");
3539 case MemOp::PortKind::Read: {
3540 // Read ports pipeline the addr and enable signals. However, the
3541 // address must be 0 for single-address memories and the enable signal
3542 // is ignored, always reading out the register. Under these constraints,
3543 // the read port can be replaced with the value from the register.
3544 replacePortField(rewriter, port, "data", memWire);
3545 break;
3546 }
3547 case MemOp::PortKind::Write: {
3548 auto data = portPipeline("data", writeStages);
3549 auto en = portPipeline("en", writeStages);
3550 auto mask = portPipeline("mask", writeStages);
3551 writes.emplace_back(data, en, mask);
3552 break;
3553 }
3554 case MemOp::PortKind::ReadWrite: {
3555 // Always read the register into the read end.
3556 replacePortField(rewriter, port, "rdata", memWire);
3557
3558 // Create a write enable and pipeline stages.
3559 auto wdata = portPipeline("wdata", writeStages);
3560 auto wmask = portPipeline("wmask", writeStages);
3561
3562 Value en = getPortFieldValue(port, "en");
3563 Value wmode = getPortFieldValue(port, "wmode");
3564
3565 auto wen = AndPrimOp::create(rewriter, port.getLoc(), en, wmode);
3566 auto wenPipelined =
3567 pipeline(wen, portClock, name + "_wen", writeStages);
3568 writes.emplace_back(wdata, wenPipelined, wmask);
3569 break;
3570 }
3571 }
3572 }
3573
3574 // Regardless of `writeUnderWrite`, always implement PortOrder.
3575 Value next = memReg;
3576 for (auto &[data, en, mask] : writes) {
3577 Value masked;
3578
3579 // If a mask bit is used, emit muxes to select the input from the
3580 // register (no mask) or the input (mask bit set).
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;
3587
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);
3593 }
3594
3595 std::reverse(chunks.begin(), chunks.end());
3596 masked = rewriter.createOrFold<CatPrimOp>(loc, chunks);
3597 next = MuxPrimOp::create(rewriter, next.getLoc(), en, masked, next);
3598 }
3599 Value typedNext = rewriter.createOrFold<BitCastOp>(next.getLoc(), ty, next);
3600 MatchingConnectOp::create(rewriter, memReg.getLoc(), memReg, typedNext);
3601
3602 // Delete the fields and their associated connects.
3603 for (Operation *conn : connects)
3604 rewriter.eraseOp(conn);
3605 for (auto portAccess : portAccesses)
3606 rewriter.eraseOp(portAccess);
3607 rewriter.eraseOp(mem);
3608
3609 return success();
3610 }
3611};
3612} // namespace
3613
3614void MemOp::getCanonicalizationPatterns(RewritePatternSet &results,
3615 MLIRContext *context) {
3616 results
3617 .insert<FoldZeroWidthMemory, FoldReadOrWriteOnlyMemory,
3618 FoldReadWritePorts, FoldUnusedPorts, FoldUnusedBits, FoldRegMems>(
3619 context);
3620}
3621
3622//===----------------------------------------------------------------------===//
3623// Declarations
3624//===----------------------------------------------------------------------===//
3625
3626// Turn synchronous reset looking register updates to registers with resets.
3627// Also, const prop registers that are driven by a mux tree containing only
3628// instances of one constant or self-assigns.
3629static LogicalResult foldHiddenReset(RegOp reg, PatternRewriter &rewriter) {
3630 // reg ; connect(reg, mux(port, const, val)) ->
3631 // reg.reset(port, const); connect(reg, val)
3632
3633 // Find the one true connect, or bail
3634 auto con = getSingleConnectUserOf(reg.getResult());
3635 if (!con)
3636 return failure();
3637
3638 auto mux = dyn_cast_or_null<MuxPrimOp>(con.getSrc().getDefiningOp());
3639 if (!mux)
3640 return failure();
3641 auto *high = mux.getHigh().getDefiningOp();
3642 auto *low = mux.getLow().getDefiningOp();
3643 // Reset value must be constant
3644 auto constOp = dyn_cast_or_null<ConstantOp>(high);
3645
3646 // Detect the case if a register only has two possible drivers:
3647 // (1) itself/uninit and (2) constant.
3648 // The mux can then be replaced with the constant.
3649 // r = mux(cond, r, 3) --> r = 3
3650 // r = mux(cond, 3, r) --> r = 3
3651 bool constReg = false;
3652
3653 if (constOp && low == reg)
3654 constReg = true;
3655 else if (dyn_cast_or_null<ConstantOp>(low) && high == reg) {
3656 constReg = true;
3657 constOp = dyn_cast<ConstantOp>(low);
3658 }
3659 if (!constOp)
3660 return failure();
3661
3662 // For a non-constant register, reset should be a module port (heuristic to
3663 // limit to intended reset lines). Replace the register anyway if constant.
3664 if (!isa<BlockArgument>(mux.getSel()) && !constReg)
3665 return failure();
3666
3667 // Check all types should be typed by now
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)
3672 return failure();
3673
3674 // Ok, we know we are doing the transformation.
3675
3676 // If we would fold the register to a constant, ensure the time-zero
3677 // simulation value is preserved.
3678 if (constReg && !preservesInitial(reg.getInitialAttr(), constOp.getValue()))
3679 return failure();
3680
3681 // Make sure the constant dominates all users.
3682 if (constOp != &con->getBlock()->front())
3683 constOp->moveBefore(&con->getBlock()->front());
3684
3685 if (!constReg) {
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);
3693 }
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);
3699 return success();
3700}
3701
3702LogicalResult RegOp::canonicalize(RegOp op, PatternRewriter &rewriter) {
3703 if (!hasDontTouch(op.getOperation()) && !op.isForceable() &&
3704 succeeded(foldHiddenReset(op, rewriter)))
3705 return success();
3706
3707 if (succeeded(demoteForceableIfUnused(op, rewriter)))
3708 return success();
3709
3710 return failure();
3711}
3712
3713//===----------------------------------------------------------------------===//
3714// Verification Ops.
3715//===----------------------------------------------------------------------===//
3716
3717static LogicalResult eraseIfZeroOrNotZero(Operation *op, Value predicate,
3718 Value enable,
3719 PatternRewriter &rewriter,
3720 bool eraseIfZero) {
3721 // If the verification op is never enabled, delete it.
3722 if (auto constant = enable.getDefiningOp<firrtl::ConstantOp>()) {
3723 if (constant.getValue().isZero()) {
3724 rewriter.eraseOp(op);
3725 return success();
3726 }
3727 }
3728
3729 // If the verification op is never triggered, delete it.
3730 if (auto constant = predicate.getDefiningOp<firrtl::ConstantOp>()) {
3731 if (constant.getValue().isZero() == eraseIfZero) {
3732 rewriter.eraseOp(op);
3733 return success();
3734 }
3735 }
3736
3737 return failure();
3738}
3739
3740template <class Op, bool EraseIfZero = false>
3741static LogicalResult canonicalizeImmediateVerifOp(Op op,
3742 PatternRewriter &rewriter) {
3743 return eraseIfZeroOrNotZero(op, op.getPredicate(), op.getEnable(), rewriter,
3744 EraseIfZero);
3745}
3746
3747void AssertOp::getCanonicalizationPatterns(RewritePatternSet &results,
3748 MLIRContext *context) {
3749 results.add(canonicalizeImmediateVerifOp<AssertOp>);
3750 results.add<patterns::AssertXWhenX>(context);
3751}
3752
3753void AssumeOp::getCanonicalizationPatterns(RewritePatternSet &results,
3754 MLIRContext *context) {
3755 results.add(canonicalizeImmediateVerifOp<AssumeOp>);
3756 results.add<patterns::AssumeXWhenX>(context);
3757}
3758
3759void UnclockedAssumeIntrinsicOp::getCanonicalizationPatterns(
3760 RewritePatternSet &results, MLIRContext *context) {
3761 results.add(canonicalizeImmediateVerifOp<UnclockedAssumeIntrinsicOp>);
3762 results.add<patterns::UnclockedAssumeIntrinsicXWhenX>(context);
3763}
3764
3765void CoverOp::getCanonicalizationPatterns(RewritePatternSet &results,
3766 MLIRContext *context) {
3767 results.add(canonicalizeImmediateVerifOp<CoverOp, /* EraseIfZero = */ true>);
3768}
3769
3770//===----------------------------------------------------------------------===//
3771// InvalidValueOp
3772//===----------------------------------------------------------------------===//
3773
3774LogicalResult InvalidValueOp::canonicalize(InvalidValueOp op,
3775 PatternRewriter &rewriter) {
3776 // Remove `InvalidValueOp`s with no uses.
3777 if (op.use_empty()) {
3778 rewriter.eraseOp(op);
3779 return success();
3780 }
3781 // Propagate invalids through a single use which is a unary op. You cannot
3782 // propagate through multiple uses as that breaks invalid semantics. Nor
3783 // can you propagate through binary ops or generally any op which computes.
3784 // Not is an exception as it is a pure, all-bits inverse.
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);
3800 return success();
3801 }
3802 return failure();
3803}
3804
3805OpFoldResult InvalidValueOp::fold(FoldAdaptor adaptor) {
3806 if (getType().getBitWidthOrSentinel() == 0 && isa<IntType>(getType()))
3807 return getIntAttr(getType(), APInt(0, 0, isa<SIntType>(getType())));
3808 return {};
3809}
3810
3811//===----------------------------------------------------------------------===//
3812// ClockGateIntrinsicOp
3813//===----------------------------------------------------------------------===//
3814
3815OpFoldResult ClockGateIntrinsicOp::fold(FoldAdaptor adaptor) {
3816 // Forward the clock if one of the enables is always true.
3817 if (isConstantOne(adaptor.getEnable()) ||
3818 isConstantOne(adaptor.getTestEnable()))
3819 return getInput();
3820
3821 // Fold to a constant zero clock if the enables are always false.
3822 if (isConstantZero(adaptor.getEnable()) &&
3823 (!getTestEnable() || isConstantZero(adaptor.getTestEnable())))
3824 return BoolAttr::get(getContext(), false);
3825
3826 // Forward constant zero clocks.
3827 if (isConstantZero(adaptor.getInput()))
3828 return BoolAttr::get(getContext(), false);
3829
3830 return {};
3831}
3832
3833LogicalResult ClockGateIntrinsicOp::canonicalize(ClockGateIntrinsicOp op,
3834 PatternRewriter &rewriter) {
3835 // Remove constant false test enable.
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(); });
3841 return success();
3842 }
3843 }
3844 }
3845
3846 return failure();
3847}
3848
3849//===----------------------------------------------------------------------===//
3850// Reference Ops.
3851//===----------------------------------------------------------------------===//
3852
3853// refresolve(forceable.ref) -> forceable.data
3854static LogicalResult
3855canonicalizeRefResolveOfForceable(RefResolveOp op, PatternRewriter &rewriter) {
3856 auto forceable = op.getRef().getDefiningOp<Forceable>();
3857 if (!forceable || !forceable.isForceable() ||
3858 op.getRef() != forceable.getDataRef() ||
3859 op.getType() != forceable.getDataType())
3860 return failure();
3861 rewriter.replaceAllUsesWith(op, forceable.getData());
3862 return success();
3863}
3864
3865void RefResolveOp::getCanonicalizationPatterns(RewritePatternSet &results,
3866 MLIRContext *context) {
3867 results.insert<patterns::RefResolveOfRefSend>(context);
3868 results.insert(canonicalizeRefResolveOfForceable);
3869}
3870
3871OpFoldResult RefCastOp::fold(FoldAdaptor adaptor) {
3872 // RefCast is unnecessary if types match.
3873 if (getInput().getType() == getType())
3874 return getInput();
3875 return {};
3876}
3877
3878static bool isConstantZero(Value operand) {
3879 auto constOp = operand.getDefiningOp<ConstantOp>();
3880 return constOp && constOp.getValue().isZero();
3881}
3882
3883template <typename Op>
3884static LogicalResult eraseIfPredFalse(Op op, PatternRewriter &rewriter) {
3885 if (isConstantZero(op.getPredicate())) {
3886 rewriter.eraseOp(op);
3887 return success();
3888 }
3889 return failure();
3890}
3891
3892void RefForceOp::getCanonicalizationPatterns(RewritePatternSet &results,
3893 MLIRContext *context) {
3894 results.add(eraseIfPredFalse<RefForceOp>);
3895}
3896void RefForceInitialOp::getCanonicalizationPatterns(RewritePatternSet &results,
3897 MLIRContext *context) {
3898 results.add(eraseIfPredFalse<RefForceInitialOp>);
3899}
3900void RefReleaseOp::getCanonicalizationPatterns(RewritePatternSet &results,
3901 MLIRContext *context) {
3902 results.add(eraseIfPredFalse<RefReleaseOp>);
3903}
3904void RefReleaseInitialOp::getCanonicalizationPatterns(
3905 RewritePatternSet &results, MLIRContext *context) {
3906 results.add(eraseIfPredFalse<RefReleaseInitialOp>);
3907}
3908
3909//===----------------------------------------------------------------------===//
3910// HasBeenResetIntrinsicOp
3911//===----------------------------------------------------------------------===//
3912
3913OpFoldResult HasBeenResetIntrinsicOp::fold(FoldAdaptor adaptor) {
3914 // The folds in here should reflect the ones for `verif::HasBeenResetOp`.
3915
3916 // Fold to zero if the reset is a constant. In this case the op is either
3917 // permanently in reset or never resets. Both mean that the reset never
3918 // finishes, so this op never returns true.
3919 if (adaptor.getReset())
3920 return getIntZerosAttr(UIntType::get(getContext(), 1));
3921
3922 // Fold to zero if the clock is a constant and the reset is synchronous. In
3923 // that case the reset will never be started.
3924 if (isUInt1(getReset().getType()) && adaptor.getClock())
3925 return getIntZerosAttr(UIntType::get(getContext(), 1));
3926
3927 return {};
3928}
3929
3930//===----------------------------------------------------------------------===//
3931// FPGAProbeIntrinsicOp
3932//===----------------------------------------------------------------------===//
3933
3934static bool isTypeEmpty(FIRRTLType type) {
3936 .Case<FVectorType>(
3937 [&](auto ty) -> bool { return isTypeEmpty(ty.getElementType()); })
3938 .Case<BundleType>([&](auto ty) -> bool {
3939 for (auto elem : ty.getElements())
3940 if (!isTypeEmpty(elem.type))
3941 return false;
3942 return true;
3943 })
3944 .Case<IntType>([&](auto ty) { return ty.getWidth() == 0; })
3945 .Default([](auto) -> bool { return false; });
3946}
3947
3948LogicalResult FPGAProbeIntrinsicOp::canonicalize(FPGAProbeIntrinsicOp op,
3949 PatternRewriter &rewriter) {
3950 auto firrtlTy = type_dyn_cast<FIRRTLType>(op.getInput().getType());
3951 if (!firrtlTy)
3952 return failure();
3953
3954 if (!isTypeEmpty(firrtlTy))
3955 return failure();
3956
3957 rewriter.eraseOp(op);
3958 return success();
3959}
3960
3961//===----------------------------------------------------------------------===//
3962// Layer Block Op
3963//===----------------------------------------------------------------------===//
3964
3965LogicalResult LayerBlockOp::canonicalize(LayerBlockOp op,
3966 PatternRewriter &rewriter) {
3967
3968 // If the layerblock is empty, erase it.
3969 if (op.getBody()->empty()) {
3970 rewriter.eraseOp(op);
3971 return success();
3972 }
3973
3974 return failure();
3975}
3976
3977//===----------------------------------------------------------------------===//
3978// Domain-related Ops
3979//===----------------------------------------------------------------------===//
3980
3981OpFoldResult UnsafeDomainCastOp::fold(FoldAdaptor adaptor) {
3982 // If no domains are specified, then forward the input to the result.
3983 if (getDomains().empty())
3984 return getInput();
3985
3986 return {};
3987}
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 &region)
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)
Definition Mem2Reg.cpp:222
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)
Definition ESIPasses.cpp:32
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.
void info(Twine message)
Definition LSPUtils.cpp:20
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.
Definition APInt.cpp:22
APInt sextZeroWidth(APInt value, unsigned width)
A safe version of APInt::sext that will NOT assert on zero-width signed APSInts.
Definition APInt.cpp:18
StringRef chooseName(StringRef a, StringRef b)
Choose a good name for an item from two options.
Definition Naming.cpp:47
static bool isConstantZero(Attribute operand)
Determine whether a constant operand is a zero value.
Definition FoldUtils.h:28
static bool isConstantOne(Attribute operand)
Determine whether a constant operand is a one value.
Definition FoldUtils.h:35
reg(value, clock, reset=None, reset_value=None, name=None, sym_name=None)
Definition seq.py:21
LogicalResult matchAndRewrite(BitsPrimOp bits, mlir::PatternRewriter &rewriter) const override