22#include "mlir/IR/Threading.h"
23#include "mlir/Pass/Pass.h"
24#include "llvm/ADT/DenseSet.h"
25#include "llvm/Support/Debug.h"
26#include "llvm/Support/LogicalResult.h"
29#define DEBUG_TYPE "firrtl-full-reset"
33#define GEN_PASS_DEF_FULLRESET
34#define GEN_PASS_DEF_MEMTOREGOFVEC
35#include "circt/Dialect/FIRRTL/Passes.h.inc"
40using namespace firrtl;
42using circt::igraph::InstanceOpInterface;
46using llvm::SmallDenseSet;
53 if (
auto arg = dyn_cast<BlockArgument>(reset)) {
54 auto module = cast<FModuleOp>(arg.getParentRegion()->getParentOp());
55 return {
module.getPortNameAttr(arg.getArgNumber()), module};
57 auto *op = reset.getDefiningOp();
58 return {op->getAttrOfType<StringAttr>(
"name"),
59 op->getParentOfType<FModuleOp>()};
88 std::optional<unsigned> existingPort;
91 ResetDomain() =
default;
94 ResetDomain(Value rootReset)
95 : rootReset(rootReset), resetName(
getResetName(rootReset)),
96 resetType(rootReset.getType()) {}
99 explicit operator bool()
const {
return static_cast<bool>(rootReset); }
103inline bool operator==(
const ResetDomain &a,
const ResetDomain &b) {
104 return (a.isTop == b.isTop && a.resetName == b.resetName &&
105 a.resetType == b.resetType);
107inline bool operator!=(
const ResetDomain &a,
const ResetDomain &b) {
116 auto it = cache.find(type);
117 if (it != cache.end())
119 auto nullBit = [&]() {
121 builder, UIntType::get(builder.getContext(), 1,
true),
126 .
Case<ClockType>([&](
auto type) {
127 return AsClockPrimOp::create(builder, nullBit());
129 .Case<AsyncResetType>([&](
auto type) {
130 return AsAsyncResetPrimOp::create(builder, nullBit());
132 .Case<SIntType, UIntType>([&](
auto type) {
133 return ConstantOp::create(
134 builder, type, APInt::getZero(type.getWidth().value_or(1)));
136 .Case<FEnumType>([&](
auto type) -> Value {
139 if (type.getNumElements() != 0 &&
140 type.getElement(0).value.getValue().isZero()) {
141 const auto &element = type.getElement(0);
143 return FEnumCreateOp::create(builder, type, element.name, value);
145 auto value = ConstantOp::create(builder,
146 UIntType::get(builder.getContext(),
149 APInt::getZero(type.getBitWidth()));
150 return BitCastOp::create(builder, type, value);
152 .Case<BundleType>([&](
auto type) {
153 auto wireOp = WireOp::create(builder, type);
154 for (
unsigned i = 0, e = type.getNumElements(); i < e; ++i) {
155 auto fieldType = type.getElementTypePreservingConst(i);
158 SubfieldOp::create(builder, fieldType, wireOp.getResult(), i);
161 return wireOp.getResult();
163 .Case<FVectorType>([&](
auto type) {
164 auto wireOp = WireOp::create(builder, type);
166 builder, type.getElementTypePreservingConst(), cache);
167 for (
unsigned i = 0, e = type.getNumElements(); i < e; ++i) {
168 auto acc = SubindexOp::create(builder, zero.getType(),
169 wireOp.getResult(), i);
172 return wireOp.getResult();
174 .Case<ResetType, AnalogType>(
175 [&](
auto type) {
return InvalidValueOp::create(builder, type); })
177 llvm_unreachable(
"switch handles all types");
180 cache.insert({type, value});
197 Value reset, Value resetValue) {
201 bool resetValueUsed =
false;
203 for (
auto &use : target.getUses()) {
204 Operation *useOp = use.getOwner();
205 builder.setInsertionPoint(useOp);
206 TypeSwitch<Operation *>(useOp)
209 .Case<ConnectOp, MatchingConnectOp>([&](
auto op) {
210 if (op.getDest() != target)
212 LLVM_DEBUG(llvm::dbgs() <<
" - Insert mux into " << op <<
"\n");
214 MuxPrimOp::create(builder, reset, resetValue, op.getSrc());
215 op.getSrcMutable().assign(muxOp);
216 resetValueUsed =
true;
219 .Case<SubfieldOp>([&](
auto op) {
221 SubfieldOp::create(builder, resetValue, op.getFieldIndexAttr());
223 resetValueUsed =
true;
225 resetSubValue.erase();
228 .Case<SubindexOp>([&](
auto op) {
230 SubindexOp::create(builder, resetValue, op.getIndexAttr());
232 resetValueUsed =
true;
234 resetSubValue.erase();
237 .Case<SubaccessOp>([&](
auto op) {
238 if (op.getInput() != target)
241 SubaccessOp::create(builder, resetValue, op.getIndex());
243 resetValueUsed =
true;
245 resetSubValue.erase();
248 return resetValueUsed;
259enum class ResetKind { Async, Sync };
261static StringRef resetKindToStringRef(
const ResetKind &kind) {
263 case ResetKind::Async:
265 case ResetKind::Sync:
268 llvm_unreachable(
"unhandled reset kind");
273struct MemToRegOfVecConverter {
274 explicit MemToRegOfVecConverter(
bool ignoreReadEnable)
275 : ignoreReadEnable(ignoreReadEnable) {}
277 void runOnModule(FModuleOp mod) {
279 mod.getBodyBlock()->walk([&](MemOp memOp) {
280 LLVM_DEBUG(llvm::dbgs() <<
"\n Memory op:" << memOp);
282 auto firMem = memOp.getSummary();
289 if (firMem.isSeqMem())
292 generateMemory(memOp, firMem);
297 Value addPipelineStages(ImplicitLocOpBuilder &b,
size_t stages, Value clock,
298 Value pipeInput, StringRef name, Value gate = {}) {
303 auto reg = RegOp::create(b, pipeInput.getType(), clock, name).getResult();
305 WhenOp::create(b, gate,
false,
306 [&]() { MatchingConnectOp::create(b, reg, pipeInput); });
308 MatchingConnectOp::create(b, reg, pipeInput);
316 Value getClock(ImplicitLocOpBuilder &builder, Value bundle) {
317 return SubfieldOp::create(builder, bundle,
"clk");
320 Value getAddr(ImplicitLocOpBuilder &builder, Value bundle) {
321 return SubfieldOp::create(builder, bundle,
"addr");
324 Value getWmode(ImplicitLocOpBuilder &builder, Value bundle) {
325 return SubfieldOp::create(builder, bundle,
"wmode");
328 Value getEnable(ImplicitLocOpBuilder &builder, Value bundle) {
329 return SubfieldOp::create(builder, bundle,
"en");
332 Value getMask(ImplicitLocOpBuilder &builder, Value bundle) {
333 auto bType = type_cast<BundleType>(bundle.getType());
334 if (bType.getElement(
"mask"))
335 return SubfieldOp::create(builder, bundle,
"mask");
336 return SubfieldOp::create(builder, bundle,
"wmask");
339 Value getData(ImplicitLocOpBuilder &builder, Value bundle,
340 bool getWdata =
false) {
341 auto bType = type_cast<BundleType>(bundle.getType());
342 if (bType.getElement(
"data"))
343 return SubfieldOp::create(builder, bundle,
"data");
344 if (bType.getElement(
"rdata") && !getWdata)
345 return SubfieldOp::create(builder, bundle,
"rdata");
346 return SubfieldOp::create(builder, bundle,
"wdata");
349 void generateRead(
const FirMemory &firMem, Value clock, Value
addr,
350 Value enable, Value
data, Value regOfVec,
351 ImplicitLocOpBuilder &builder) {
352 if (ignoreReadEnable) {
355 for (
size_t j = 0, e = firMem.
readLatency; j != e; ++j) {
356 auto enLast = enable;
358 enable = addPipelineStages(builder, 1, clock, enable,
"en");
359 addr = addPipelineStages(builder, 1, clock,
addr,
"addr", enLast);
365 addPipelineStages(builder, firMem.
readLatency, clock, enable,
"en");
371 Value
rdata = SubaccessOp::create(builder, regOfVec,
addr);
372 if (!ignoreReadEnable) {
374 MatchingConnectOp::create(
375 builder,
data, InvalidValueOp::create(builder,
data.getType()));
377 WhenOp::create(builder, enable,
false, [&]() {
378 MatchingConnectOp::create(builder,
data,
rdata);
382 MatchingConnectOp::create(builder,
data,
rdata);
386 void generateWrite(
const FirMemory &firMem, Value clock, Value
addr,
387 Value enable, Value maskBits, Value wdataIn,
388 Value regOfVec, ImplicitLocOpBuilder &builder) {
393 addr = addPipelineStages(builder, numStages, clock,
addr,
"addr");
394 enable = addPipelineStages(builder, numStages, clock, enable,
"en");
395 wdataIn = addPipelineStages(builder, numStages, clock, wdataIn,
"wdata");
396 maskBits = addPipelineStages(builder, numStages, clock, maskBits,
"wmask");
405 SmallVector<std::tuple<Value, Value, Value>, 8> loweredRegDataMaskFields;
420 if (!getFields(
rdata, wdataIn, maskBits, loweredRegDataMaskFields,
422 wdataIn.getDefiningOp()->emitOpError(
423 "Cannot convert memory to bank of registers");
427 WhenOp::create(builder, enable,
false, [&]() {
429 for (
auto regDataMask : loweredRegDataMaskFields) {
430 auto regField = std::get<0>(regDataMask);
431 auto dataField = std::get<1>(regDataMask);
432 auto maskField = std::get<2>(regDataMask);
434 WhenOp::create(builder, maskField,
false, [&]() {
435 MatchingConnectOp::create(builder, regField, dataField);
441 void generateReadWrite(
const FirMemory &firMem, Value clock, Value
addr,
442 Value enable, Value maskBits, Value wdataIn,
443 Value rdataOut, Value
wmode, Value regOfVec,
444 ImplicitLocOpBuilder &builder) {
449 addr = addPipelineStages(builder, numStages, clock,
addr,
"addr");
450 enable = addPipelineStages(builder, numStages, clock, enable,
"en");
451 wdataIn = addPipelineStages(builder, numStages, clock, wdataIn,
"wdata");
452 maskBits = addPipelineStages(builder, numStages, clock, maskBits,
"wmask");
455 Value
rdata = SubaccessOp::create(builder, regOfVec,
addr);
457 SmallVector<std::tuple<Value, Value, Value>, 8> loweredRegDataMaskFields;
458 if (!getFields(
rdata, wdataIn, maskBits, loweredRegDataMaskFields,
460 wdataIn.getDefiningOp()->emitOpError(
461 "Cannot convert memory to bank of registers");
465 MatchingConnectOp::create(
466 builder, rdataOut, InvalidValueOp::create(builder, rdataOut.getType()));
468 WhenOp::create(builder, enable,
false, [&]() {
471 builder,
wmode,
true,
475 for (
auto regDataMask : loweredRegDataMaskFields) {
476 auto regField = std::get<0>(regDataMask);
477 auto dataField = std::get<1>(regDataMask);
478 auto maskField = std::get<2>(regDataMask);
481 builder, maskField,
false, [&]() {
482 MatchingConnectOp::create(builder, regField, dataField);
487 [&]() { MatchingConnectOp::create(builder, rdataOut,
rdata); });
498 bool getFields(Value reg, Value input, Value
mask,
499 SmallVectorImpl<std::tuple<Value, Value, Value>> &results,
500 ImplicitLocOpBuilder &builder) {
504 if (
auto bundle = type_dyn_cast<BundleType>(inType)) {
505 if (
auto mBundle = type_dyn_cast<BundleType>(maskType))
506 return mBundle.getNumElements() == bundle.getNumElements();
507 }
else if (
auto vec = type_dyn_cast<FVectorType>(inType)) {
508 if (
auto mVec = type_dyn_cast<FVectorType>(maskType))
509 return mVec.getNumElements() == vec.getNumElements();
515 std::function<bool(Value, Value, Value)> flatAccess =
516 [&](Value
reg, Value input, Value
mask) ->
bool {
517 FIRRTLType inType = type_cast<FIRRTLType>(input.getType());
518 if (!isValidMask(inType, type_cast<FIRRTLType>(
mask.getType()))) {
519 input.getDefiningOp()->emitOpError(
"Mask type is not valid");
523 .
Case<BundleType>([&](BundleType bundle) {
524 for (
size_t i = 0, e = bundle.getNumElements(); i != e; ++i) {
525 auto regField = SubfieldOp::create(builder, reg, i);
526 auto inputField = SubfieldOp::create(builder, input, i);
527 auto maskField = SubfieldOp::create(builder,
mask, i);
528 if (!flatAccess(regField, inputField, maskField))
533 .Case<FVectorType>([&](
auto vector) {
534 for (
size_t i = 0, e = vector.getNumElements(); i != e; ++i) {
535 auto regField = SubindexOp::create(builder, reg, i);
536 auto inputField = SubindexOp::create(builder, input, i);
537 auto maskField = SubindexOp::create(builder,
mask, i);
538 if (!flatAccess(regField, inputField, maskField))
543 .Case<IntType>([&](
auto iType) {
544 results.push_back({
reg, input,
mask});
545 return iType.getWidth().has_value();
547 .Default([&](
auto) {
return false; });
549 if (flatAccess(reg, input,
mask))
555 void generateMemory(MemOp memOp,
FirMemory &firMem) {
556 ImplicitLocOpBuilder builder(memOp.getLoc(), memOp);
557 auto dataType = memOp.getDataType();
559 auto innerSym = memOp.getInnerSym();
560 SmallVector<Value> debugPorts;
563 for (
size_t index = 0, rend = memOp.getNumResults(); index < rend;
565 auto result = memOp.getResult(index);
566 if (type_isa<RefType>(result.getType())) {
567 debugPorts.push_back(result);
572 auto wire = WireOp::create(
573 builder, result.getType(),
574 (memOp.getName() +
"_" + memOp.getPortName(index)).str(),
575 memOp.getNameKind());
576 result.replaceAllUsesWith(wire.getResult());
577 result = wire.getResult();
579 auto adr = getAddr(builder, result);
580 auto enb = getEnable(builder, result);
581 auto clk = getClock(builder, result);
582 auto dta = getData(builder, result);
587 RegOp::create(builder, FVectorType::get(dataType, firMem.
depth),
588 clk, memOp.getNameAttr());
591 if (!memOp.getAnnotationsAttr().empty())
592 regOfVec.setAnnotationsAttr(memOp.getAnnotationsAttr());
594 regOfVec.setInnerSymAttr(memOp.getInnerSymAttr());
596 auto portKind = memOp.getPortKind(index);
597 if (portKind == MemOp::PortKind::Read) {
598 generateRead(firMem,
clk, adr, enb, dta, regOfVec.getResult(), builder);
599 }
else if (portKind == MemOp::PortKind::Write) {
600 auto mask = getMask(builder, result);
601 generateWrite(firMem,
clk, adr, enb,
mask, dta, regOfVec.getResult(),
604 auto wmode = getWmode(builder, result);
605 auto wDta = getData(builder, result,
true);
606 auto mask = getMask(builder, result);
607 generateReadWrite(firMem,
clk, adr, enb,
mask, wDta, dta,
wmode,
608 regOfVec.getResult(), builder);
615 for (
auto r : debugPorts)
616 r.replaceAllUsesWith(RefSendOp::create(builder, regOfVec.getResult()));
619 bool ignoreReadEnable =
false;
620 unsigned numConverted = 0;
625 unsigned &numConverted) {
626 MemToRegOfVecConverter converter(ignoreReadEnable);
627 converter.runOnModule(mod);
628 numConverted += converter.numConverted;
632struct MemToRegOfVecPass
633 :
public circt::firrtl::impl::MemToRegOfVecBase<MemToRegOfVecPass> {
636 void runOnOperation()
override {
637 auto circtOp = getOperation();
638 auto &instanceInfo = getAnalysis<InstanceInfo>();
641 convertMemToRegOfVecAnnoClass))
642 return markAllAnalysesPreserved();
644 SmallVector<FModuleOp> modules;
645 for (
auto moduleOp : circtOp.getOps<FModuleOp>())
646 if (instanceInfo.anyInstanceInEffectiveDesign(moduleOp))
647 modules.push_back(moduleOp);
649 std::atomic<unsigned> totalConverted{0};
650 mlir::parallelForEach(&getContext(), modules, [&](FModuleOp mod) {
651 unsigned numConverted = 0;
653 totalConverted += numConverted;
655 numConvertedMems += totalConverted.load();
661struct FullResetRunner {
664 InstanceInfo &instanceInfo,
bool convertAsyncDomainMems)
665 : circuit(circuit), instanceGraph(&ig),
666 instancePathCache(&instancePathCache), instanceInfo(&instanceInfo),
667 convertAsyncDomainMems(convertAsyncDomainMems) {}
674 LogicalResult collectAnnos();
680 FailureOr<std::optional<Value>> collectAnnos(FModuleOp module);
682 LogicalResult buildDomains();
683 void buildDomains(FModuleOp module,
const InstancePath &instPath,
685 unsigned indent = 0);
687 void convertMemsInAsyncDomains();
689 LogicalResult determineImpl();
690 LogicalResult determineImpl(FModuleOp module, ResetDomain &domain);
692 LogicalResult implementFullReset();
693 LogicalResult implementFullReset(FModuleOp module, ResetDomain &domain);
694 LogicalResult implementFullReset(Operation *op, FModuleOp module,
698 LogicalResult implementFullReset(FInstanceLike inst, StringAttr moduleName,
706 DenseMap<Operation *, Value> annotatedResets;
721 bool convertAsyncDomainMems =
false;
725LogicalResult FullResetRunner::run() {
726 if (failed(collectAnnos()))
728 if (failed(buildDomains()))
730 if (convertAsyncDomainMems)
731 convertMemsInAsyncDomains();
732 if (failed(determineImpl()))
734 if (failed(implementFullReset()))
739void FullResetRunner::convertMemsInAsyncDomains() {
740 SmallVector<FModuleOp> asyncDomainMods;
741 for (
auto &[mod, entries] : domains) {
744 auto &domain = entries.back().first;
745 if (!domain.rootReset)
747 if (!type_isa<AsyncResetType>(domain.resetType))
749 if (!instanceInfo->anyInstanceInEffectiveDesign(mod))
751 asyncDomainMods.push_back(mod);
753 if (asyncDomainMods.empty())
757 llvm::dbgs() <<
"\n";
758 debugHeader(
"Convert comb mems in async full-reset domains") <<
"\n\n";
759 for (
auto mod : asyncDomainMods)
763 mlir::parallelForEach(
764 circuit.getContext(), asyncDomainMods, [&](FModuleOp mod) {
765 unsigned converted = 0;
766 runCombMemsToRegOfVec(mod, false, converted);
772 bool convertAsyncDomainMems) {
774 return FullResetRunner(circuit, ig, instancePathCache, instanceInfo,
775 convertAsyncDomainMems)
783LogicalResult FullResetRunner::collectAnnos() {
785 llvm::dbgs() <<
"\n";
788 SmallVector<std::pair<FModuleOp, std::optional<Value>>> results;
789 for (
auto module : circuit.getOps<FModuleOp>())
790 results.push_back({module, {}});
792 if (failed(mlir::failableParallelForEach(
793 circuit.getContext(), results, [&](
auto &moduleAndResult) {
794 auto result = collectAnnos(moduleAndResult.first);
797 moduleAndResult.second = *result;
802 for (
auto [module, reset] : results)
803 if (reset.has_value())
804 annotatedResets.insert({module, *reset});
808FailureOr<std::optional<Value>>
809FullResetRunner::collectAnnos(FModuleOp module) {
810 bool anyFailed =
false;
817 if (anno.
isClass(excludeFromFullResetAnnoClass)) {
819 conflictingAnnos.insert({anno, module.getLoc()});
822 if (anno.
isClass(fullResetAnnoClass)) {
824 module.emitError("''FullResetAnnotation' cannot target module; must
"
825 "target port or wire/node instead
");
833 // Consume any reset annotations on module ports.
835 // Helper for checking annotations and determining the reset
836 auto checkAnnotations = [&](Annotation anno, Value arg) {
837 if (anno.isClass(fullResetAnnoClass)) {
838 ResetKind expectedResetKind;
839 if (auto rt = anno.getMember<StringAttr>("resetType
")) {
841 expectedResetKind = ResetKind::Sync;
842 } else if (rt == "async
") {
843 expectedResetKind = ResetKind::Async;
845 mlir::emitError(arg.getLoc(),
846 "'FullResetAnnotation' requires resetType ==
'sync' "
847 "|
'async', but got resetType ==
")
853 mlir::emitError(arg.getLoc(),
854 "'FullResetAnnotation' requires resetType ==
"
855 "'sync' |
'async', but got no resetType
");
859 // Check that the type is well-formed
860 bool isAsync = expectedResetKind == ResetKind::Async;
861 bool validUint = false;
862 if (auto uintT = dyn_cast<UIntType>(arg.getType()))
863 validUint = uintT.getWidth() == 1;
864 if ((isAsync && !isa<AsyncResetType>(arg.getType())) ||
865 (!isAsync && !validUint)) {
866 auto kind = resetKindToStringRef(expectedResetKind);
867 mlir::emitError(arg.getLoc(),
868 "'FullResetAnnotation' with resetType ==
'")
869 << kind << "' must target
" << kind << " reset, but targets
"
876 conflictingAnnos.insert({anno, reset.getLoc()});
880 if (anno.isClass(excludeFromFullResetAnnoClass)) {
882 mlir::emitError(arg.getLoc(),
883 "'ExcludeFromFullResetAnnotation' cannot
"
884 "target port/wire/node; must target
module instead");
892 Value arg =
module.getArgument(argNum);
893 return checkAnnotations(anno, arg);
899 module.getBody().walk([&](Operation *op) {
901 if (!isa<WireOp, NodeOp>(op)) {
902 if (AnnotationSet::hasAnnotation(op, fullResetAnnoClass,
903 excludeFromFullResetAnnoClass)) {
906 "reset annotations must target module, port, or wire/node");
914 auto arg = op->getResult(0);
915 return checkAnnotations(anno, arg);
924 if (!ignore && !reset) {
925 LLVM_DEBUG(llvm::dbgs()
926 <<
"No reset annotation for " << module.getName() <<
"\n");
927 return std::optional<Value>();
931 if (conflictingAnnos.size() > 1) {
932 auto diag =
module.emitError("multiple reset annotations on module '")
933 << module.getName() << "'";
934 for (
auto &annoAndLoc : conflictingAnnos)
935 diag.attachNote(annoAndLoc.second)
936 <<
"conflicting " << annoAndLoc.first.getClassAttr() <<
":";
942 llvm::dbgs() <<
"Annotated reset for " <<
module.getName() << ": ";
944 llvm::dbgs() <<
"no domain\n";
945 else if (
auto arg = dyn_cast<BlockArgument>(reset))
946 llvm::dbgs() <<
"port " <<
module.getPortName(arg.getArgNumber()) << "\n";
948 llvm::dbgs() <<
"wire "
949 << reset.getDefiningOp()->getAttrOfType<StringAttr>(
"name")
955 return std::optional<Value>(reset);
967LogicalResult FullResetRunner::buildDomains() {
969 llvm::dbgs() <<
"\n";
974 auto &instGraph = *instanceGraph;
980 dyn_cast_or_null<FModuleOp>(node.
getModule().getOperation()))
981 buildDomains(module,
InstancePath{}, Value{}, instGraph);
985 bool anyFailed =
false;
986 for (
auto &it : domains) {
987 auto module = cast<FModuleOp>(it.first);
988 auto &domainConflicts = it.second;
989 if (domainConflicts.size() <= 1)
993 SmallDenseSet<Value> printedDomainResets;
994 auto diag =
module.emitError("module '")
996 << "' instantiated in different reset domains";
997 for (
auto &it : domainConflicts) {
998 ResetDomain &domain = it.first;
999 const auto &path = it.second;
1000 auto inst = path.leaf();
1001 auto loc = path.empty() ?
module.getLoc() : inst.getLoc();
1002 auto ¬e = diag.attachNote(loc);
1006 note <<
"root instance";
1008 note <<
"instance '";
1011 [&](InstanceOpInterface inst) { note << inst.getInstanceName(); },
1012 [&]() { note <<
"/"; });
1018 if (domain.rootReset) {
1020 note <<
" reset domain rooted at '" << nameAndModule.first.getValue()
1021 <<
"' of module '" << nameAndModule.second.getName() <<
"'";
1024 if (printedDomainResets.insert(domain.rootReset).second) {
1025 diag.attachNote(domain.rootReset.getLoc())
1026 <<
"reset domain '" << nameAndModule.first.getValue()
1027 <<
"' of module '" << nameAndModule.second.getName()
1028 <<
"' declared here:";
1031 note <<
" no reset domain";
1034 return failure(anyFailed);
1037void FullResetRunner::buildDomains(FModuleOp module,
1042 llvm::dbgs().indent(indent * 2) <<
"Visiting ";
1043 if (instPath.
empty())
1044 llvm::dbgs() <<
"$root";
1046 llvm::dbgs() << instPath.
leaf().getInstanceName();
1047 llvm::dbgs() <<
" (" <<
module.getName() << ")\n";
1052 auto it = annotatedResets.find(module);
1053 if (it != annotatedResets.end()) {
1056 if (
auto localReset = it->second)
1057 domain = ResetDomain(localReset);
1058 domain.isTop =
true;
1059 }
else if (parentReset) {
1061 domain = ResetDomain(parentReset);
1068 auto &entries = domains[module];
1069 if (domain.rootReset)
1070 if (llvm::all_of(entries,
1071 [&](
const auto &entry) {
return entry.first != domain; }))
1072 entries.push_back({domain, instPath});
1075 for (
auto *record : *instGraph[module]) {
1076 auto submodule = dyn_cast<FModuleOp>(*record->getTarget()->getModule());
1080 instancePathCache->appendInstance(instPath, record->getInstance());
1081 buildDomains(submodule, childPath, domain.rootReset, instGraph, indent + 1);
1086LogicalResult FullResetRunner::determineImpl() {
1087 auto anyFailed =
false;
1089 llvm::dbgs() <<
"\n";
1090 debugHeader(
"Determine implementation") <<
"\n\n";
1092 for (
auto &it : domains) {
1093 auto module = cast<FModuleOp>(it.first);
1094 auto &entries = it.second;
1096 if (entries.empty())
1098 auto &domain = entries.back().first;
1099 if (failed(determineImpl(module, domain)))
1102 return failure(anyFailed);
1120LogicalResult FullResetRunner::determineImpl(FModuleOp module,
1121 ResetDomain &domain) {
1125 LLVM_DEBUG(llvm::dbgs() <<
"Planning reset for " << module.getName() <<
"\n");
1130 LLVM_DEBUG(llvm::dbgs()
1131 <<
"- Rooting at local value " << domain.resetName <<
"\n");
1132 domain.localReset = domain.rootReset;
1133 if (
auto blockArg = dyn_cast<BlockArgument>(domain.rootReset))
1134 domain.existingPort = blockArg.getArgNumber();
1140 auto neededName = domain.resetName;
1141 auto neededType = domain.resetType;
1142 LLVM_DEBUG(llvm::dbgs() <<
"- Looking for existing port " << neededName
1144 auto portNames =
module.getPortNames();
1145 auto *portIt = llvm::find(portNames, neededName);
1148 if (portIt == portNames.end()) {
1149 LLVM_DEBUG(llvm::dbgs() <<
"- Creating new port " << neededName <<
"\n");
1150 domain.resetName = neededName;
1154 LLVM_DEBUG(llvm::dbgs() <<
"- Reusing existing port " << neededName <<
"\n");
1157 auto portNo = std::distance(portNames.begin(), portIt);
1158 auto portType =
module.getPortType(portNo);
1159 if (portType != neededType) {
1160 auto diag = emitError(module.getPortLocation(portNo),
"module '")
1161 <<
module.getName() << "' is in reset domain requiring port '"
1162 << domain.resetName.getValue() << "' to have type "
1163 << domain.resetType << ", but has type " << portType;
1164 diag.attachNote(domain.rootReset.getLoc()) <<
"reset domain rooted here";
1169 domain.existingPort = portNo;
1170 domain.localReset =
module.getArgument(portNo);
1179LogicalResult FullResetRunner::implementFullReset() {
1181 llvm::dbgs() <<
"\n";
1184 for (
auto &it : domains) {
1185 auto module = cast<FModuleOp>(it.first);
1186 auto &entries = it.second;
1190 if (!entries.empty())
1191 domain = entries.back().first;
1192 if (failed(implementFullReset(module, domain)))
1203LogicalResult FullResetRunner::implementFullReset(FModuleOp module,
1204 ResetDomain &domain) {
1209 SmallVector<FInstanceLike> instances;
1210 module.walk([&](FInstanceLike instOp) { instances.push_back(instOp); });
1212 if (!instances.empty())
1213 llvm::dbgs() <<
"Tie off instances in " << module.getName() <<
"\n";
1215 for (
auto instOp : instances)
1216 if (failed(implementFullReset(instOp, module, Value())))
1221 LLVM_DEBUG(llvm::dbgs() <<
"Implementing full reset for " << module.getName()
1225 auto *
context =
module.getContext();
1227 annotations.addAnnotations(DictionaryAttr::get(
1229 StringAttr::get(
context, fullResetAnnoClass))));
1230 annotations.applyToOperation(module);
1233 auto actualReset = domain.localReset;
1234 if (!domain.localReset) {
1235 PortInfo portInfo{domain.resetName,
1239 domain.rootReset.getLoc()};
1240 module.insertPorts({{0, portInfo}});
1241 actualReset =
module.getArgument(0);
1242 LLVM_DEBUG(llvm::dbgs() <<
"- Inserted port " << domain.resetName <<
"\n");
1246 llvm::dbgs() <<
"- Using ";
1247 if (
auto blockArg = dyn_cast<BlockArgument>(actualReset))
1248 llvm::dbgs() <<
"port #" << blockArg.getArgNumber() <<
" ";
1250 llvm::dbgs() <<
"wire/node ";
1256 SmallVector<Operation *> opsToUpdate;
1257 module.walk([&](Operation *op) {
1258 if (isa<FInstanceLike, RegOp, RegResetOp>(op))
1259 opsToUpdate.push_back(op);
1266 if (!isa<BlockArgument>(actualReset)) {
1267 mlir::DominanceInfo dom(module);
1272 auto *resetOp = actualReset.getDefiningOp();
1273 if (!opsToUpdate.empty() && !dom.dominates(resetOp, opsToUpdate[0])) {
1274 LLVM_DEBUG(llvm::dbgs()
1275 <<
"- Reset doesn't dominate all uses, needs to be moved\n");
1279 auto nodeOp = dyn_cast<NodeOp>(resetOp);
1280 if (nodeOp && !dom.dominates(nodeOp.getInput(), opsToUpdate[0])) {
1281 LLVM_DEBUG(llvm::dbgs()
1282 <<
"- Promoting node to wire for move: " << nodeOp <<
"\n");
1283 auto builder = ImplicitLocOpBuilder::atBlockBegin(nodeOp.getLoc(),
1284 nodeOp->getBlock());
1285 auto wireOp = WireOp::create(
1286 builder, nodeOp.getResult().getType(), nodeOp.getNameAttr(),
1287 nodeOp.getNameKindAttr(), nodeOp.getAnnotationsAttr(),
1288 nodeOp.getInnerSymAttr(), nodeOp.getForceableAttr());
1290 nodeOp->replaceAllUsesWith(wireOp);
1291 nodeOp->removeAttr(nodeOp.getInnerSymAttrName());
1295 nodeOp.setNameKind(NameKindEnum::DroppableName);
1296 nodeOp.setAnnotationsAttr(ArrayAttr::get(builder.getContext(), {}));
1297 builder.setInsertionPointAfter(nodeOp);
1298 emitConnect(builder, wireOp.getResult(), nodeOp.getResult());
1300 actualReset = wireOp.getResult();
1301 domain.localReset = wireOp.getResult();
1306 Block *targetBlock = dom.findNearestCommonDominator(
1307 resetOp->getBlock(), opsToUpdate[0]->getBlock());
1309 if (targetBlock != resetOp->getBlock())
1310 llvm::dbgs() <<
"- Needs to be moved to different block\n";
1319 auto getParentInBlock = [](Operation *op,
Block *block) {
1320 while (op && op->getBlock() != block)
1321 op = op->getParentOp();
1324 auto *resetOpInTarget = getParentInBlock(resetOp, targetBlock);
1325 auto *firstOpInTarget = getParentInBlock(opsToUpdate[0], targetBlock);
1331 if (resetOpInTarget->isBeforeInBlock(firstOpInTarget))
1332 resetOp->moveBefore(resetOpInTarget);
1334 resetOp->moveBefore(firstOpInTarget);
1339 for (
auto *op : opsToUpdate)
1340 if (failed(implementFullReset(op, module, actualReset)))
1348LogicalResult FullResetRunner::implementFullReset(FInstanceLike inst,
1349 StringAttr moduleName,
1350 Value actualReset) {
1354 auto *node = instanceGraph->lookup(moduleName);
1355 auto refModule = dyn_cast<FModuleOp>(*node->
getModule());
1358 auto *domainIt = domains.find(refModule);
1359 if (domainIt == domains.end() || domainIt->second.empty())
1361 auto &domain = domainIt->second.back().first;
1362 assert(domain &&
"null domains should not be listed");
1364 ImplicitLocOpBuilder builder(inst.getLoc(), inst);
1366 LLVM_DEBUG(llvm::dbgs() << (actualReset ?
"- Update " :
"- Tie-off ")
1372 if (!domain.localReset) {
1373 LLVM_DEBUG(llvm::dbgs() <<
" - Adding new result as reset\n");
1374 auto newInstOp = inst.cloneWithInsertedPortsAndReplaceUses(
1376 {domain.resetName, domain.resetType, Direction::In}}});
1377 instReset = newInstOp->getResult(0);
1378 instanceGraph->replaceInstance(inst, newInstOp);
1381 }
else if (domain.existingPort.has_value()) {
1382 auto idx = *domain.existingPort;
1383 instReset = inst->getResult(idx);
1384 LLVM_DEBUG(llvm::dbgs() <<
" - Using result #" << idx <<
" as reset\n");
1393 builder.setInsertionPointAfter(inst);
1400 LLVM_DEBUG(llvm::dbgs() <<
" - Tying off reset to constant 0\n");
1401 if (type_isa<AsyncResetType>(domain.resetType))
1402 actualReset = SpecialConstantOp::create(builder, domain.resetType,
false);
1404 actualReset = ConstantOp::create(
1405 builder, UIntType::get(builder.getContext(), 1), APInt(1, 0));
1409 assert(instReset && actualReset);
1417LogicalResult FullResetRunner::implementFullReset(Operation *op,
1419 Value actualReset) {
1420 ImplicitLocOpBuilder builder(op->getLoc(), op);
1423 if (
auto instOp = dyn_cast<FInstanceLike>(op))
1424 return implementFullReset(
1425 instOp, cast<StringAttr>(instOp.getReferencedModuleNamesAttr()[0]),
1433 if (
auto regOp = dyn_cast<RegOp>(op)) {
1434 LLVM_DEBUG(llvm::dbgs() <<
"- Adding full reset to " << regOp <<
"\n");
1436 auto newRegOp = RegResetOp::create(
1437 builder, regOp.getResult().getType(), regOp.getClockVal(), actualReset,
1438 zero, regOp.getNameAttr(), regOp.getNameKindAttr(),
1439 regOp.getAnnotations(), regOp.getInnerSymAttr(),
1440 regOp.getForceableAttr());
1441 regOp.getResult().replaceAllUsesWith(newRegOp.getResult());
1442 if (regOp.getForceable())
1443 regOp.getRef().replaceAllUsesWith(newRegOp.getRef());
1449 if (
auto regOp = dyn_cast<RegResetOp>(op)) {
1452 if (type_isa<AsyncResetType>(regOp.getResetSignal().getType()) ||
1453 type_isa<UIntType>(actualReset.getType())) {
1454 LLVM_DEBUG(llvm::dbgs() <<
"- Skipping (has reset) " << regOp <<
"\n");
1457 if (failed(regOp.verifyInvariants()))
1461 LLVM_DEBUG(llvm::dbgs() <<
"- Updating reset of " << regOp <<
"\n");
1463 auto reset = regOp.getResetSignal();
1464 auto value = regOp.getResetValue();
1470 builder.setInsertionPointAfterValue(regOp.getResult());
1471 auto mux = MuxPrimOp::create(builder, reset, value, regOp.getResult());
1475 builder.setInsertionPoint(regOp);
1477 regOp.getResetSignalMutable().assign(actualReset);
1478 regOp.getResetValueMutable().assign(zero);
1485 :
public circt::firrtl::impl::FullResetBase<FullResetPass> {
1486 using FullResetBase::FullResetBase;
1488 void runOnOperation()
override {
1489 auto &ig = getAnalysis<InstanceGraph>();
1490 auto &instanceInfo = getAnalysis<InstanceInfo>();
1491 if (failed(
runFullReset(getOperation(), ig, instanceInfo,
1493 return signalPassFailure();
1494 markAnalysesPreserved<InstanceGraph, InstanceInfo>();
assert(baseType &&"element must be base type")
static std::unique_ptr< Context > context
static Value createZeroValue(ImplicitLocOpBuilder &builder, FIRRTLBaseType type, SmallDenseMap< FIRRTLBaseType, Value > &cache)
Construct a zero value of the given type using the given builder.
static StringAttr getResetName(Value reset)
Return the name of a reset.
static bool insertResetMux(ImplicitLocOpBuilder &builder, Value target, Value reset, Value resetValue)
Helper function that inserts reset multiplexer into all ConnectOps with the given target.
static std::pair< StringAttr, FModuleOp > getResetNameAndModule(Value reset)
Return the name and parent module of a reset.
This class provides a read-only projection over the MLIR attributes that represent a set of annotatio...
bool removeAnnotations(llvm::function_ref< bool(Annotation)> predicate)
Remove all annotations from this annotation set for which predicate returns true.
static bool removePortAnnotations(Operation *module, llvm::function_ref< bool(unsigned, Annotation)> predicate)
Remove all port annotations from a module or extmodule for which predicate returns true.
This class provides a read-only projection of an annotation.
bool isClass(Args... names) const
Return true if this annotation matches any of the specified class names.
FIRRTLBaseType getConstType(bool isConst) const
Return a 'const' or non-'const' version of this type.
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 graph tracks modules and where they are instantiated.
This is a Node in the InstanceGraph.
bool noUses()
Return true if there are no more instances of this module.
auto getModule()
Get the module that this node is tracking.
An instance path composed of a series of instances.
InstanceOpInterface leaf() const
std::string getInstanceName(mlir::func::CallOp callOp)
A helper function to get the instance name.
mlir::TypedValue< FIRRTLBaseType > FIRRTLBaseValue
void runCombMemsToRegOfVec(FModuleOp mod, bool ignoreReadEnable, unsigned &numConverted)
LogicalResult runFullReset(CircuitOp circuit, InstanceGraph &ig, InstanceInfo &instanceInfo, bool convertAsyncDomainMems=false)
void emitConnect(OpBuilder &builder, Location loc, Value lhs, Value rhs)
Emit a connect between two values.
StringAttr getName(ArrayAttr names, size_t idx)
Return the name at the specified index of the ArrayAttr or null if it cannot be determined.
static bool operator==(const ModulePort &a, const ModulePort &b)
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.
llvm::raw_ostream & debugHeader(const llvm::Twine &str, unsigned width=80)
Write a "header"-like string to the debug stream with a certain width.
bool operator!=(uint64_t a, const FVInt &b)
int run(Type[Generator] generator=CppGenerator, List[str] cmdline_args=sys.argv)
reg(value, clock, reset=None, reset_value=None, name=None, sym_name=None)
This holds the name and type that describes the module's ports.
A data structure that caches and provides paths to module instances in the IR.