18#include "mlir/Analysis/TopologicalSortUtils.h"
19#include "mlir/Dialect/Arith/IR/Arith.h"
20#include "mlir/IR/Builders.h"
21#include "mlir/IR/DialectImplementation.h"
22#include "mlir/IR/Matchers.h"
23#include "mlir/IR/PatternMatch.h"
26#include "llvm/ADT/SmallString.h"
33 auto memType = cast<seq::HLMemType>(hlmemHandle.getType());
34 auto shape = memType.getShape();
35 if (shape.size() != addresses.size())
38 for (
auto [dim, addr] : llvm::zip(shape, addresses)) {
39 auto addrType = dyn_cast<IntegerType>(addr.getType());
42 if (addrType.getIntOrFloatBitWidth() != llvm::Log2_64_Ceil(dim))
51 if (result.attributes.getNamed(
"name"))
55 StringRef resultName = parser.getResultName(0).first;
56 if (!resultName.empty() &&
isdigit(resultName[0]))
58 result.addAttribute(
"name", parser.getBuilder().getStringAttr(resultName));
62 if (!op->hasAttr(
"name"))
65 auto name = op->getAttrOfType<StringAttr>(
"name").getValue();
69 SmallString<32> resultNameStr;
70 llvm::raw_svector_ostream tmpStream(resultNameStr);
71 p.printOperand(op->getResult(0), tmpStream);
72 auto actualName = tmpStream.str().drop_front();
73 return actualName == name;
78 std::optional<OpAsmParser::UnresolvedOperand> operand,
86 Value operand, Type type) {
91 OpAsmParser &parser, Type refType,
92 std::optional<OpAsmParser::UnresolvedOperand> operand, Type &type) {
94 type = seq::ImmutableType::get(refType);
99 Type refType, Value operand,
108ParseResult ReadPortOp::parse(OpAsmParser &parser, OperationState &result) {
109 llvm::SMLoc loc = parser.getCurrentLocation();
111 OpAsmParser::UnresolvedOperand memOperand, rdenOperand;
112 bool hasRdEn =
false;
113 llvm::SmallVector<OpAsmParser::UnresolvedOperand, 2> addressOperands;
114 seq::HLMemType memType;
116 if (parser.parseOperand(memOperand) ||
117 parser.parseOperandList(addressOperands, OpAsmParser::Delimiter::Square))
120 if (succeeded(parser.parseOptionalKeyword(
"rden"))) {
121 if (failed(parser.parseOperand(rdenOperand)))
126 if (parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() ||
127 parser.parseType(memType))
130 llvm::SmallVector<Type> operandTypes = memType.getAddressTypes();
131 operandTypes.insert(operandTypes.begin(), memType);
133 llvm::SmallVector<OpAsmParser::UnresolvedOperand> allOperands = {memOperand};
134 llvm::copy(addressOperands, std::back_inserter(allOperands));
136 operandTypes.push_back(parser.getBuilder().getI1Type());
137 allOperands.push_back(rdenOperand);
140 if (parser.resolveOperands(allOperands, operandTypes, loc, result.operands))
143 result.addTypes(memType.getElementType());
145 llvm::SmallVector<int32_t, 2> operandSizes;
146 operandSizes.push_back(1);
147 operandSizes.push_back(addressOperands.size());
148 operandSizes.push_back(hasRdEn ? 1 : 0);
149 result.addAttribute(
"operandSegmentSizes",
150 parser.getBuilder().getDenseI32ArrayAttr(operandSizes));
154void ReadPortOp::print(OpAsmPrinter &p) {
155 p <<
" " << getMemory() <<
"[" << getAddresses() <<
"]";
157 p <<
" rden " << getRdEn();
158 p.printOptionalAttrDict((*this)->getAttrs(), {
"operandSegmentSizes"});
159 p <<
" : " << getMemory().getType();
163 auto memName = getMemory().getDefiningOp<seq::HLMemOp>().
getName();
164 setNameFn(getReadData(), (memName +
"_rdata").str());
167void ReadPortOp::build(OpBuilder &builder, OperationState &result, Value memory,
168 ValueRange addresses, Value rdEn,
unsigned latency) {
169 auto memType = cast<seq::HLMemType>(memory.getType());
170 ReadPortOp::build(builder, result, memType.getElementType(), memory,
171 addresses, rdEn, latency);
178ParseResult WritePortOp::parse(OpAsmParser &parser, OperationState &result) {
179 llvm::SMLoc loc = parser.getCurrentLocation();
180 OpAsmParser::UnresolvedOperand memOperand, dataOperand, wrenOperand;
181 llvm::SmallVector<OpAsmParser::UnresolvedOperand, 2> addressOperands;
182 seq::HLMemType memType;
184 if (parser.parseOperand(memOperand) ||
185 parser.parseOperandList(addressOperands,
186 OpAsmParser::Delimiter::Square) ||
187 parser.parseOperand(dataOperand) || parser.parseKeyword(
"wren") ||
188 parser.parseOperand(wrenOperand) ||
189 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() ||
190 parser.parseType(memType))
193 llvm::SmallVector<Type> operandTypes = memType.getAddressTypes();
194 operandTypes.insert(operandTypes.begin(), memType);
195 operandTypes.push_back(memType.getElementType());
196 operandTypes.push_back(parser.getBuilder().getI1Type());
198 llvm::SmallVector<OpAsmParser::UnresolvedOperand, 2> allOperands(
200 allOperands.insert(allOperands.begin(), memOperand);
201 allOperands.push_back(dataOperand);
202 allOperands.push_back(wrenOperand);
204 if (parser.resolveOperands(allOperands, operandTypes, loc, result.operands))
210void WritePortOp::print(OpAsmPrinter &p) {
211 p <<
" " << getMemory() <<
"[" << getAddresses() <<
"] " << getInData()
212 <<
" wren " << getWrEn();
213 p.printOptionalAttrDict((*this)->getAttrs());
214 p <<
" : " << getMemory().getType();
222 setNameFn(getHandle(),
getName());
225void HLMemOp::build(OpBuilder &builder, OperationState &result, Value clk,
226 Value rst, StringRef name, llvm::ArrayRef<int64_t> shape,
228 HLMemType t = HLMemType::get(builder.getContext(), shape,
elementType);
229 HLMemOp::build(builder, result, t, clk, rst, name);
238 IntegerAttr &threshold,
239 Type &outputFlagType,
240 StringRef directive) {
242 if (succeeded(parser.parseOptionalKeyword(directive))) {
243 int64_t thresholdValue;
244 if (succeeded(parser.parseInteger(thresholdValue))) {
245 threshold = parser.getBuilder().getI64IntegerAttr(thresholdValue);
246 outputFlagType = parser.getBuilder().getI1Type();
249 return parser.emitError(parser.getNameLoc(),
250 "expected integer value after " + directive +
257 Type &outputFlagType) {
263 Type &outputFlagType) {
269 Type outputFlagType) {
272 <<
" " << threshold.getInt();
276 Type outputFlagType) {
279 <<
" " << threshold.getInt();
283 setNameFn(getOutput(),
"out");
284 setNameFn(getEmpty(),
"empty");
285 setNameFn(getFull(),
"full");
286 if (
auto ae = getAlmostEmpty())
287 setNameFn(ae,
"almostEmpty");
288 if (
auto af = getAlmostFull())
289 setNameFn(af,
"almostFull");
292LogicalResult FIFOOp::verify() {
293 auto aet = getAlmostEmptyThreshold();
294 auto aft = getAlmostFullThreshold();
295 size_t depth = getDepth();
296 if (aft.has_value() && aft.value() > depth)
297 return emitOpError(
"almost full threshold must be <= FIFO depth");
299 if (aet.has_value() && aet.value() > depth)
300 return emitOpError(
"almost empty threshold must be <= FIFO depth");
313 setNameFn(getResult(), *name);
316template <
typename TOp>
318 if ((op.getReset() ==
nullptr) ^ (op.getResetValue() ==
nullptr))
319 return op->emitOpError(
320 "either reset and resetValue or neither must be specified");
321 bool hasReset = op.getReset() !=
nullptr;
322 if (hasReset && op.getResetValue().getType() != op.getInput().getType())
323 return op->emitOpError(
"reset value must be the same type as the input");
328std::optional<size_t> CompRegOp::getTargetResultIndex() {
return 0; }
330LogicalResult CompRegOp::verify() {
return verifyResets(*
this); }
340 setNameFn(getResult(), *name);
343std::optional<size_t> CompRegClockEnabledOp::getTargetResultIndex() {
347LogicalResult CompRegClockEnabledOp::verify() {
return verifyResets(*
this); }
350 PatternRewriter &rewriter) {
353 auto *inputOp = op.getInput().getDefiningOp();
354 if (isa_and_nonnull<comb::MuxOp, arith::SelectOp>(inputOp) &&
355 inputOp->getOperand(0) == op.getClockEnable()) {
356 rewriter.modifyOpInPlace(
357 op, [&] { op.getInputMutable().assign(inputOp->getOperand(1)); });
363 if (mlir::matchPattern(op.getClockEnable(), mlir::m_ConstantInt(&en))) {
364 if (
en.isAllOnes()) {
366 op, op.getInput(), op.getClk(), op.getNameAttr(), op.getReset(),
367 op.getResetValue(), op.getInitialValue(), op.getInnerSymAttr());
382 setNameFn(getResult(), *name);
385std::optional<size_t> ShiftRegOp::getTargetResultIndex() {
return 0; }
387LogicalResult ShiftRegOp::verify() {
397void FirRegOp::build(OpBuilder &builder, OperationState &result, Value input,
398 Value clk, StringAttr name, hw::InnerSymAttr innerSym,
401 OpBuilder::InsertionGuard guard(builder);
403 result.addOperands(input);
404 result.addOperands(clk);
406 result.addAttribute(getNameAttrName(result.name), name);
409 result.addAttribute(getInnerSymAttrName(result.name), innerSym);
412 result.addAttribute(getPresetAttrName(result.name), preset);
414 result.addTypes(input.getType());
417void FirRegOp::build(OpBuilder &builder, OperationState &result, Value input,
418 Value clk, StringAttr name, Value reset, Value resetValue,
419 hw::InnerSymAttr innerSym,
bool isAsync,
422 OpBuilder::InsertionGuard guard(builder);
424 result.addOperands(input);
425 result.addOperands(clk);
426 result.addOperands(reset);
427 result.addOperands(resetValue);
429 result.addAttribute(getNameAttrName(result.name), name);
431 result.addAttribute(getIsAsyncAttrName(result.name), builder.getUnitAttr());
434 result.addAttribute(getInnerSymAttrName(result.name), innerSym);
437 result.addAttribute(getPresetAttrName(result.name), preset);
439 result.addTypes(input.getType());
442ParseResult FirRegOp::parse(OpAsmParser &parser, OperationState &result) {
443 auto &builder = parser.getBuilder();
444 llvm::SMLoc loc = parser.getCurrentLocation();
446 using Op = OpAsmParser::UnresolvedOperand;
449 if (parser.parseOperand(next) || parser.parseKeyword(
"clock") ||
450 parser.parseOperand(clk))
453 if (succeeded(parser.parseOptionalKeyword(
"sym"))) {
454 hw::InnerSymAttr innerSym;
455 if (parser.parseCustomAttributeWithFallback(innerSym,
nullptr,
456 "inner_sym", result.attributes))
461 std::optional<std::pair<Op, Op>> resetAndValue;
462 if (succeeded(parser.parseOptionalKeyword(
"reset"))) {
464 if (succeeded(parser.parseOptionalKeyword(
"async")))
466 else if (succeeded(parser.parseOptionalKeyword(
"sync")))
469 return parser.emitError(loc,
"invalid reset, expected 'sync' or 'async'");
471 result.attributes.append(
"isAsync", builder.getUnitAttr());
473 resetAndValue = {{}, {}};
474 if (parser.parseOperand(resetAndValue->first) || parser.parseComma() ||
475 parser.parseOperand(resetAndValue->second))
479 std::optional<APInt> presetValue;
480 llvm::SMLoc presetValueLoc;
481 if (succeeded(parser.parseOptionalKeyword(
"preset"))) {
482 presetValueLoc = parser.getCurrentLocation();
483 OptionalParseResult presetIntResult =
484 parser.parseOptionalInteger(presetValue.emplace());
485 if (!presetIntResult.has_value() || failed(*presetIntResult))
486 return parser.emitError(presetValueLoc,
"expected integer value");
487 if (presetValue->isNegative())
488 return parser.emitError(presetValueLoc,
489 "preset value must not be negative");
493 if (parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() ||
494 parser.parseType(ty))
496 result.addTypes({ty});
500 if (hw::type_isa<seq::ClockType>(ty)) {
505 return parser.emitError(presetValueLoc,
506 "cannot preset register of unknown width");
510 APInt presetResult = presetValue->sextOrTrunc(width);
511 if (presetResult.zextOrTrunc(presetValue->getBitWidth()) != *presetValue)
512 return parser.emitError(presetValueLoc,
"preset value too large");
514 auto builder = parser.getBuilder();
515 auto presetTy = builder.getIntegerType(width);
516 auto resultAttr = builder.getIntegerAttr(presetTy, presetResult);
517 result.addAttribute(
"preset", resultAttr);
522 if (parser.resolveOperand(next, ty, result.operands))
525 Type clkTy = ClockType::get(result.getContext());
526 if (parser.resolveOperand(clk, clkTy, result.operands))
530 Type i1 = IntegerType::get(result.getContext(), 1);
531 if (parser.resolveOperand(resetAndValue->first, i1, result.operands) ||
532 parser.resolveOperand(resetAndValue->second, ty, result.operands))
539void FirRegOp::print(::mlir::OpAsmPrinter &p) {
540 SmallVector<StringRef> elidedAttrs = {
541 getInnerSymAttrName(), getIsAsyncAttrName(), getPresetAttrName()};
543 p <<
' ' << getNext() <<
" clock " << getClk();
545 if (
auto sym = getInnerSymAttr()) {
551 p <<
" reset " << (getIsAsync() ?
"async" :
"sync") <<
' ';
552 p << getReset() <<
", " << getResetValue();
555 if (
auto preset = getPresetAttr()) {
559 const auto &presetVal = preset.getValue();
560 if (presetVal.isNonNegative())
563 p << presetVal.zext(presetVal.getBitWidth() + 1);
567 elidedAttrs.push_back(
"name");
569 p.printOptionalAttrDict((*this)->getAttrs(), elidedAttrs);
570 p <<
" : " << getNext().getType();
574LogicalResult FirRegOp::verify() {
575 if (getReset() || getResetValue() || getIsAsync()) {
576 if (!getReset() || !getResetValue())
577 return emitOpError(
"must specify reset and reset value");
580 return emitOpError(
"register with no reset cannot be async");
582 if (
auto preset = getPresetAttr()) {
585 if (preset.getType() != getType() && presetWidth != width)
586 return emitOpError(
"preset type width must match register type");
596 setNameFn(getResult(),
getName());
599std::optional<size_t> FirRegOp::getTargetResultIndex() {
return 0; }
601LogicalResult FirRegOp::canonicalize(FirRegOp op, PatternRewriter &rewriter) {
605 if (
auto reset = op.getReset()) {
607 if (constOp.getValue().isZero()) {
608 rewriter.replaceOpWithNewOp<FirRegOp>(
609 op, op.getNext(), op.getClk(), op.getNameAttr(),
610 op.getInnerSymAttr(), op.getPresetAttr());
617 if (op.getInnerSymAttr())
625 if (op.getNext() == op.getResult())
627 if (
auto clk = op.getClk().getDefiningOp<seq::ToClockOp>())
633 bool replaceWithConstZero =
true;
634 if (
auto preset = op.getPresetAttr())
635 if (!preset.getValue().isZero())
636 replaceWithConstZero =
false;
638 if (
isConstant() && !op.getResetValue() && replaceWithConstZero) {
639 if (isa<seq::ClockType>(op.getType())) {
640 rewriter.replaceOpWithNewOp<seq::ConstClockOp>(
641 op, seq::ClockConstAttr::get(rewriter.getContext(), ClockConst::Low));
645 rewriter.replaceOpWithNewOp<
hw::BitcastOp>(op, op.getType(), constant);
655 if (
auto nextMux = op.getNext().getDefiningOp<
comb::MuxOp>()) {
657 if (op.getPresetAttr())
664 if (nextMux.getTrueValue() == op.getResult() &&
665 matchPattern(nextMux.getFalseValue(), m_Constant(&value))) {
666 replacedValue = nextMux.getFalseValue();
669 else if (nextMux.getFalseValue() == op.getResult() &&
670 matchPattern(nextMux.getTrueValue(), m_Constant(&value))) {
671 replacedValue = nextMux.getTrueValue();
679 if (op.getResetValue()) {
680 Attribute resetConst;
681 if (matchPattern(op.getResetValue(), m_Constant(&resetConst))) {
682 if (resetConst != value)
692 rewriter.replaceOp(op, replacedValue);
702 if (!op.getReset() && !op.getPresetAttr()) {
706 if (isa<IntegerType>(
707 hw::type_cast<hw::ArrayType>(op.getResult().getType())
708 .getElementType())) {
709 SmallVector<Value> nextOperands;
710 bool changed =
false;
711 for (
const auto &[i, value] :
712 llvm::enumerate(arrayCreate.getOperands())) {
713 auto index = arrayCreate.getOperands().size() - i - 1;
717 if (arrayGet.getInput() == op.getResult() &&
718 matchPattern(arrayGet.getIndex(),
719 m_ConstantInt(&elementIndex)) &&
720 elementIndex == index) {
722 rewriter, op.getLoc(),
727 nextOperands.push_back(value);
732 rewriter, arrayCreate.getLoc(), nextOperands);
733 if (arrayCreate->hasOneUse())
736 rewriter.replaceOp(arrayCreate, newNextVal);
739 rewriter.replaceOpWithNewOp<FirRegOp>(op, newNextVal, op.getClk(),
741 op.getInnerSymAttr());
753OpFoldResult FirRegOp::fold(FoldAdaptor adaptor) {
756 if (getInnerSymAttr())
759 auto presetAttr = getPresetAttr();
769 if (
auto reset = getReset())
771 if (constOp.getValue().isOne())
772 return getResetValue();
777 bool isTrivialFeedback = (getNext() == getResult());
778 bool isNeverClocked =
779 adaptor.getClk() !=
nullptr;
780 if (!isTrivialFeedback && !isNeverClocked)
785 if (
auto resetValue = getResetValue()) {
786 if (
auto *op = resetValue.getDefiningOp()) {
787 if (op->hasTrait<OpTrait::ConstantLike>() && !presetAttr)
789 if (
auto constOp = dyn_cast<hw::ConstantOp>(op))
790 if (presetAttr.getValue() == constOp.getValue())
798 auto intType = dyn_cast<IntegerType>(getType());
804 return IntegerAttr::get(intType, 0);
811OpFoldResult ClockGateOp::fold(FoldAdaptor adaptor) {
820 return ClockConstAttr::get(getContext(), ClockConst::Low);
823 if (
auto clockAttr = dyn_cast_or_null<ClockConstAttr>(adaptor.getInput()))
824 if (clockAttr.getValue() == ClockConst::Low)
825 return ClockConstAttr::get(getContext(), ClockConst::Low);
829 auto clockGateInputOp = getInput().getDefiningOp<ClockGateOp>();
830 while (clockGateInputOp) {
831 if (clockGateInputOp.getEnable() == getEnable() &&
832 clockGateInputOp.getTestEnable() == getTestEnable())
834 clockGateInputOp = clockGateInputOp.getInput().getDefiningOp<ClockGateOp>();
840LogicalResult ClockGateOp::canonicalize(ClockGateOp op,
841 PatternRewriter &rewriter) {
843 if (
auto testEnable = op.getTestEnable()) {
845 if (constOp.getValue().isZero()) {
846 rewriter.modifyOpInPlace(op,
847 [&] { op.getTestEnableMutable().clear(); });
856std::optional<size_t> ClockGateOp::getTargetResultIndex() {
864OpFoldResult ClockMuxOp::fold(FoldAdaptor adaptor) {
866 return getTrueClock();
868 return getFalseClock();
876LogicalResult ClockDividerOp::canonicalize(ClockDividerOp op,
877 PatternRewriter &rewriter) {
879 if (
auto innerDiv = op.getInput().getDefiningOp<ClockDividerOp>()) {
880 auto outerPow2 = op.getPow2();
881 auto innerPow2 = innerDiv.getPow2();
882 auto combinedPow2 = outerPow2 + innerPow2;
884 rewriter.replaceOpWithNewOp<ClockDividerOp>(op, innerDiv.getInput(),
895LogicalResult FirMemOp::canonicalize(FirMemOp op, PatternRewriter &rewriter) {
897 if (op.getInnerSymAttr())
900 bool readOnly =
true, writeOnly =
true;
903 for (
auto *user : op->getUsers()) {
904 if (isa<FirMemReadOp, FirMemReadWriteOp>(user)) {
907 if (isa<FirMemWriteOp, FirMemReadWriteOp>(user)) {
910 assert((isa<FirMemReadOp, FirMemWriteOp, FirMemReadWriteOp>(user)) &&
911 "invalid seq.firmem user");
914 for (
auto *user :
llvm::make_early_inc_range(op->getUsers()))
915 rewriter.eraseOp(user);
917 rewriter.eraseOp(op);
921 if (readOnly && !op.getInit()) {
923 for (
auto *user :
llvm::make_early_inc_range(op->getUsers())) {
924 auto readOp = cast<FirMemReadOp>(user);
926 rewriter, readOp.getLoc(),
928 if (readOp.getType() != zero.getType())
930 readOp.getType(), zero);
931 rewriter.replaceOp(readOp, zero);
933 rewriter.eraseOp(op);
940 auto nameAttr = (*this)->getAttrOfType<StringAttr>(
"name");
941 if (!nameAttr.getValue().empty())
942 setNameFn(getResult(), nameAttr.getValue());
945std::optional<size_t> FirMemOp::getTargetResultIndex() {
return 0; }
949 if (
auto mask = op.getMask()) {
950 auto memType = op.getMemory().getType();
951 if (!memType.getMaskWidth())
952 return op.emitOpError(
"has mask operand but memory type '")
953 << memType <<
"' has no mask";
954 auto expected = IntegerType::get(op.getContext(), *memType.getMaskWidth());
955 if (mask.getType() != expected)
956 return op.emitOpError(
"has mask operand of type '")
957 << mask.getType() <<
"', but memory type requires '" << expected
964LogicalResult FirMemReadWriteOp::verify() {
return verifyFirMemMask(*
this); }
969 return value.getDefiningOp<seq::ConstClockOp>();
975 return constOp.getValue().isZero();
982 return constOp.getValue().isAllOnes();
986LogicalResult FirMemReadOp::canonicalize(FirMemReadOp op,
987 PatternRewriter &rewriter) {
990 rewriter.modifyOpInPlace(op, [&] { op.getEnableMutable().erase(0); });
996LogicalResult FirMemWriteOp::canonicalize(FirMemWriteOp op,
997 PatternRewriter &rewriter) {
1001 auto memOp = op.getMemory().getDefiningOp<FirMemOp>();
1002 if (memOp.getInnerSymAttr())
1004 rewriter.eraseOp(op);
1007 bool anyChanges =
false;
1011 rewriter.modifyOpInPlace(op, [&] { op.getEnableMutable().erase(0); });
1017 rewriter.modifyOpInPlace(op, [&] { op.getMaskMutable().erase(0); });
1021 return success(anyChanges);
1024LogicalResult FirMemReadWriteOp::canonicalize(FirMemReadWriteOp op,
1025 PatternRewriter &rewriter) {
1030 auto opAttrs = op->getAttrs();
1031 auto opAttrNames = op.getAttributeNames();
1032 auto newOp = rewriter.replaceOpWithNewOp<FirMemReadOp>(
1033 op, op.getMemory(), op.getAddress(), op.getClk(), op.getEnable());
1034 for (
auto namedAttr : opAttrs)
1035 if (!
llvm::is_contained(opAttrNames, namedAttr.
getName()))
1036 newOp->setAttr(namedAttr.
getName(), namedAttr.getValue());
1039 bool anyChanges =
false;
1043 rewriter.modifyOpInPlace(op, [&] { op.getEnableMutable().erase(0); });
1049 rewriter.modifyOpInPlace(op, [&] { op.getMaskMutable().erase(0); });
1053 return success(anyChanges);
1060OpFoldResult ConstClockOp::fold(FoldAdaptor adaptor) {
1061 return ClockConstAttr::get(getContext(), getValue());
1068LogicalResult ToClockOp::canonicalize(ToClockOp op, PatternRewriter &rewriter) {
1069 if (
auto fromClock = op.getInput().getDefiningOp<FromClockOp>()) {
1070 rewriter.replaceOp(op, fromClock.getInput());
1076OpFoldResult ToClockOp::fold(FoldAdaptor adaptor) {
1077 if (
auto fromClock = getInput().getDefiningOp<FromClockOp>())
1078 return fromClock.getInput();
1079 if (
auto intAttr = dyn_cast_or_null<IntegerAttr>(adaptor.getInput())) {
1081 intAttr.getValue().isZero() ? ClockConst::Low : ClockConst::High;
1082 return ClockConstAttr::get(getContext(), value);
1087LogicalResult FromClockOp::canonicalize(FromClockOp op,
1088 PatternRewriter &rewriter) {
1089 if (
auto toClock = op.getInput().getDefiningOp<ToClockOp>()) {
1090 rewriter.replaceOp(op, toClock.getInput());
1096OpFoldResult FromClockOp::fold(FoldAdaptor adaptor) {
1097 if (
auto toClock = getInput().getDefiningOp<ToClockOp>())
1098 return toClock.getInput();
1099 if (
auto clockAttr = dyn_cast_or_null<ClockConstAttr>(adaptor.getInput())) {
1100 auto ty = IntegerType::get(getContext(), 1);
1101 return IntegerAttr::get(ty, clockAttr.getValue() == ClockConst::High);
1110OpFoldResult ClockInverterOp::fold(FoldAdaptor adaptor) {
1111 if (
auto chainedInv = getInput().getDefiningOp<ClockInverterOp>())
1112 return chainedInv.getInput();
1113 if (
auto clockAttr = dyn_cast_or_null<ClockConstAttr>(adaptor.getInput())) {
1114 auto clockIn = clockAttr.getValue() == ClockConst::High;
1115 return ClockConstAttr::get(getContext(),
1116 clockIn ? ClockConst::Low : ClockConst::High);
1126 depth = op->getAttrOfType<IntegerAttr>(
"depth").
getInt();
1127 numReadPorts = op->getAttrOfType<IntegerAttr>(
"numReadPorts").getUInt();
1128 numWritePorts = op->getAttrOfType<IntegerAttr>(
"numWritePorts").getUInt();
1130 op->getAttrOfType<IntegerAttr>(
"numReadWritePorts").getUInt();
1131 readLatency = op->getAttrOfType<IntegerAttr>(
"readLatency").getUInt();
1132 writeLatency = op->getAttrOfType<IntegerAttr>(
"writeLatency").getUInt();
1133 dataWidth = op->getAttrOfType<IntegerAttr>(
"width").getUInt();
1134 if (op->hasAttrOfType<IntegerAttr>(
"maskGran"))
1135 maskGran = op->getAttrOfType<IntegerAttr>(
"maskGran").getUInt();
1137 maskGran = dataWidth;
1138 readUnderWrite = op->getAttrOfType<seq::RUWAttr>(
"readUnderWrite").getValue();
1140 op->getAttrOfType<seq::WUWAttr>(
"writeUnderWrite").getValue();
1141 if (
auto clockIDsAttr = op->getAttrOfType<ArrayAttr>(
"writeClockIDs"))
1142 for (
auto clockID : clockIDsAttr)
1143 writeClockIDs.push_back(
1144 cast<IntegerAttr>(clockID).getValue().getZExtValue());
1145 initFilename = op->getAttrOfType<StringAttr>(
"initFilename").getValue();
1146 initIsBinary = op->getAttrOfType<BoolAttr>(
"initIsBinary").getValue();
1147 initIsInline = op->getAttrOfType<BoolAttr>(
"initIsInline").getValue();
1150LogicalResult InitialOp::verify() {
1152 auto *terminator = this->getBody().front().getTerminator();
1153 if (terminator->getOperands().size() != getNumResults())
1154 return emitError() <<
"result type doesn't match with the terminator";
1155 for (
auto [lhs, rhs] :
1156 llvm::zip(terminator->getOperands().getTypes(), getResultTypes())) {
1157 if (cast<seq::ImmutableType>(rhs).getInnerType() != lhs)
1158 return emitError() << cast<seq::ImmutableType>(rhs).getInnerType()
1159 <<
" is expected but got " << lhs;
1162 auto blockArgs = this->getBody().front().getArguments();
1164 if (blockArgs.size() != getNumOperands())
1165 return emitError() <<
"operand type doesn't match with the block arg";
1167 for (
auto [blockArg, operand] :
llvm::zip(blockArgs, getOperands())) {
1168 if (blockArg.getType() !=
1169 cast<ImmutableType>(operand.getType()).getInnerType())
1171 << blockArg.getType() <<
" is expected but got "
1172 << cast<ImmutableType>(operand.getType()).getInnerType();
1176void InitialOp::build(OpBuilder &builder, OperationState &result,
1177 TypeRange resultTypes, std::function<
void()> ctor) {
1178 OpBuilder::InsertionGuard guard(builder);
1180 builder.createBlock(result.addRegion());
1181 SmallVector<Type> types;
1182 for (
auto t : resultTypes)
1183 types.push_back(
seq::ImmutableType::
get(t));
1185 result.addTypes(types);
1191TypedValue<seq::ImmutableType>
1193 mlir::IntegerAttr attr) {
1194 auto initial = seq::InitialOp::create(builder, loc, attr.getType(), [&]() {
1195 auto constant = hw::ConstantOp::create(builder, loc, attr);
1196 seq::YieldOp::create(builder, loc, ArrayRef<Value>{constant});
1198 return cast<TypedValue<seq::ImmutableType>>(initial->getResult(0));
1201mlir::TypedValue<seq::ImmutableType>
1203 assert(op->getNumResults() == 1 &&
1204 op->hasTrait<mlir::OpTrait::ConstantLike>());
1205 auto initial = seq::InitialOp::create(
1206 builder, op->getLoc(), op->getResultTypes(), [&]() {
1207 auto clonedOp = builder.clone(*op);
1208 seq::YieldOp::create(builder, op->getLoc(), clonedOp->getResults());
1210 return cast<mlir::TypedValue<seq::ImmutableType>>(initial.getResult(0));
1214 auto resultNum = cast<OpResult>(value).getResultNumber();
1215 auto initialOp = value.getDefiningOp<seq::InitialOp>();
1217 return initialOp.getBodyBlock()->getTerminator()->getOperand(resultNum);
1221 SmallVector<Operation *> initialOps;
1222 for (
auto &op : *block)
1223 if (isa<seq::InitialOp>(op))
1224 initialOps.push_back(&op);
1226 if (!mlir::computeTopologicalSorting(initialOps, {}))
1227 return block->getParentOp()->emitError() <<
"initial ops cannot be "
1228 <<
"topologically sorted";
1231 if (initialOps.size() <= 1)
1232 return initialOps.empty() ? seq::InitialOp()
1233 : cast<seq::InitialOp>(initialOps[0]);
1235 auto initialOp = cast<seq::InitialOp>(initialOps.front());
1236 auto yieldOp = cast<seq::YieldOp>(initialOp.getBodyBlock()->getTerminator());
1239 resultToYieldOperand;
1241 for (
auto [result, operand] :
1242 llvm::zip(initialOp.getResults(), yieldOp->getOperands()))
1243 resultToYieldOperand.insert({result, operand});
1245 for (
size_t i = 1; i < initialOps.size(); ++i) {
1246 auto currentInitialOp = cast<seq::InitialOp>(initialOps[i]);
1247 auto operands = currentInitialOp->getOperands();
1248 for (
auto [blockArg, operand] :
1250 if (
auto initOp = operand.getDefiningOp<seq::InitialOp>()) {
1251 assert(resultToYieldOperand.count(operand) &&
1252 "it must be visited already");
1253 blockArg.replaceAllUsesWith(resultToYieldOperand.lookup(operand));
1256 initialOp.getBodyBlock()->addArgument(
1257 cast<seq::ImmutableType>(operand.getType()).getInnerType(),
1259 initialOp.getInputsMutable().append(operand);
1263 auto currentYieldOp =
1264 cast<seq::YieldOp>(currentInitialOp.getBodyBlock()->getTerminator());
1266 for (
auto [result, operand] :
llvm::zip(currentInitialOp.getResults(),
1267 currentYieldOp->getOperands()))
1268 resultToYieldOperand.insert({result, operand});
1271 yieldOp.getOperandsMutable().append(currentYieldOp.getOperands());
1272 currentYieldOp->erase();
1276 initialOp.getBodyBlock()->getOperations().splice(
1277 initialOp.end(), currentInitialOp.getBodyBlock()->getOperations());
1281 yieldOp->moveBefore(initialOp.getBodyBlock(),
1282 initialOp.getBodyBlock()->end());
1284 auto builder = OpBuilder::atBlockBegin(block);
1285 SmallVector<Type> types;
1286 for (
auto [result, operand] : resultToYieldOperand)
1287 types.push_back(operand.getType());
1291 auto newInitial = seq::InitialOp::create(builder, initialOp.getLoc(), types);
1292 newInitial.getInputsMutable().append(initialOp.getInputs());
1294 for (
auto [resultAndOperand, newResult] :
1295 llvm::zip(resultToYieldOperand, newInitial.getResults()))
1296 resultAndOperand.first.replaceAllUsesWith(newResult);
1299 for (
auto oldBlockArg : initialOp.
getBodyBlock()->getArguments()) {
1300 auto blockArg = newInitial.getBodyBlock()->addArgument(
1301 oldBlockArg.getType(), oldBlockArg.getLoc());
1302 oldBlockArg.replaceAllUsesWith(blockArg);
1305 newInitial.getBodyBlock()->getOperations().splice(
1306 newInitial.end(), initialOp.getBodyBlock()->getOperations());
1309 while (!initialOps.empty())
1310 initialOps.pop_back_val()->erase();
1320#define GET_OP_CLASSES
1321#include "circt/Dialect/Seq/Seq.cpp.inc"
assert(baseType &&"element must be base type")
static bool isConstZero(Value value)
static std::optional< APInt > getInt(Value value)
Helper to convert a value to a constant integer if it is one.
static Block * getBodyBlock(FModuleLike mod)
void printFIFOAFThreshold(OpAsmPrinter &p, Operation *op, IntegerAttr threshold, Type outputFlagType)
static bool isConstClock(Value value)
static ParseResult parseFIFOFlagThreshold(OpAsmParser &parser, IntegerAttr &threshold, Type &outputFlagType, StringRef directive)
static void printOptionalTypeMatch(OpAsmPrinter &p, Operation *op, Type refType, Value operand, Type type)
static bool isConstAllOnes(Value value)
static ParseResult parseOptionalImmutableTypeMatch(OpAsmParser &parser, Type refType, std::optional< OpAsmParser::UnresolvedOperand > operand, Type &type)
void printFIFOAEThreshold(OpAsmPrinter &p, Operation *op, IntegerAttr threshold, Type outputFlagType)
LogicalResult verifyResets(TOp op)
static bool canElideName(OpAsmPrinter &p, Operation *op)
ParseResult parseFIFOAEThreshold(OpAsmParser &parser, IntegerAttr &threshold, Type &outputFlagType)
static LogicalResult verifyFirMemMask(Op op)
static void printOptionalImmutableTypeMatch(OpAsmPrinter &p, Operation *op, Type refType, Value operand, Type type)
static ParseResult parseOptionalTypeMatch(OpAsmParser &parser, Type refType, std::optional< OpAsmParser::UnresolvedOperand > operand, Type &type)
static void setNameFromResult(OpAsmParser &parser, OperationState &result)
ParseResult parseFIFOAFThreshold(OpAsmParser &parser, IntegerAttr &threshold, Type &outputFlagType)
static InstancePath empty
create(elements, Type result_type=None)
Direction get(bool isOutput)
Returns an output direction if isOutput is true, otherwise returns an input direction.
bool isConstant(Operation *op)
Return true if the specified operation has a constant value.
StringAttr getName(ArrayAttr names, size_t idx)
Return the name at the specified index of the ArrayAttr or null if it cannot be determined.
int64_t getBitWidth(mlir::Type type)
Return the hardware bit width of a type.
FailureOr< seq::InitialOp > mergeInitialOps(Block *block)
bool isValidIndexValues(Value hlmemHandle, ValueRange addresses)
mlir::TypedValue< seq::ImmutableType > createConstantInitialValue(OpBuilder builder, Location loc, mlir::IntegerAttr attr)
Value unwrapImmutableValue(mlir::TypedValue< seq::ImmutableType > immutableVal)
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.
static bool isConstantZero(Attribute operand)
Determine whether a constant operand is a zero value.
static bool isConstantOne(Attribute operand)
Determine whether a constant operand is a one value.
function_ref< void(Value, StringRef)> OpAsmSetValueNameFn
FirMemory(hw::HWModuleGeneratedOp op)