17#include "mlir/IR/PatternMatch.h"
18#include "mlir/Pass/Pass.h"
19#include "mlir/Transforms/DialectConversion.h"
20#include "mlir/Transforms/RegionUtils.h"
21#include "llvm/Support/Debug.h"
23#define DEBUG_TYPE "convert-to-arcs"
30using mlir::ConversionConfig;
35 return op->hasTrait<OpTrait::ConstantLike>() ||
37 ClockedOpInterface, seq::InitialOp, seq::ClockGateOp,
38 sim::DPICallOp, llhd::ProbeOp>(op) ||
39 op->getNumResults() > 1 || op->getNumRegions() > 0 ||
40 !mlir::isMemoryEffectFree(op);
44 SmallVectorImpl<Value> &values) {
45 if (!reg.getInitialValue())
46 return values.push_back({}), success();
50 OpBuilder builder(reg);
51 auto init = seq::FromImmutableOp::create(builder, reg.getLoc(), reg.getType(),
52 reg.getInitialValue());
54 values.push_back(init);
64 LogicalResult
run(ModuleOp module);
66 LogicalResult analyzeFanIn();
74 SmallVector<Operation *> arcBreakers;
78 SmallVector<Operation *> postOrder;
88 SmallVector<mlir::CallOpInterface> arcUses;
96LogicalResult Converter::run(ModuleOp module) {
97 for (
auto &op : module.getOps())
98 if (auto sym = dyn_cast<
mlir::SymbolOpInterface>(&op))
99 globalNamespace.newName(sym.
getName());
100 for (
auto module : module.getOps<
HWModuleOp>())
101 if (failed(runOnModule(module)))
106LogicalResult Converter::runOnModule(
HWModuleOp module) {
109 arcBreakerIndices.clear();
111 if (isa<seq::InitialOp>(&op))
115 arcBreakerIndices[&op] = arcBreakers.size();
116 arcBreakers.push_back(&op);
119 if (module.getBodyBlock()->without_terminator().empty() &&
120 isa<hw::OutputOp>(module.getBodyBlock()->getTerminator()))
122 LLVM_DEBUG(llvm::dbgs() <<
"Analyzing " << module.getModuleNameAttr() <<
" ("
123 << arcBreakers.size() <<
" breakers)\n");
128 if (failed(analyzeFanIn()))
134 if (failed(absorbRegs(module)))
140LogicalResult Converter::analyzeFanIn() {
141 SmallVector<std::tuple<Operation *, SmallVector<Value, 2>>> worklist;
142 SetVector<Value> seenOperands;
143 auto addToWorklist = [&](Operation *op) {
144 seenOperands.clear();
145 for (
auto operand : op->getOperands())
146 seenOperands.insert(operand);
147 mlir::getUsedValuesDefinedAbove(op->getRegions(), seenOperands);
148 worklist.emplace_back(op, seenOperands.getArrayRef());
153 for (
auto *op : arcBreakers) {
154 unsigned index = arcBreakerIndices.lookup(op);
155 auto mask = APInt::getOneBitSet(arcBreakers.size(), index);
156 faninMasks[op] =
mask;
161 DenseSet<Operation *> seen;
162 DenseSet<Operation *> finished;
164 while (!worklist.empty()) {
165 auto &[op, operands] = worklist.back();
166 if (operands.empty()) {
168 postOrder.push_back(op);
174 auto operand = operands.pop_back_val();
175 auto *definingOp = operand.getDefiningOp();
177 finished.contains(definingOp))
179 if (!seen.insert(definingOp).second) {
180 definingOp->emitError(
"combinational loop detected");
183 addToWorklist(definingOp);
185 LLVM_DEBUG(llvm::dbgs() <<
"- Sorted " << postOrder.size() <<
" ops\n");
191 for (
auto *op :
llvm::reverse(postOrder)) {
192 auto mask = APInt::getZero(arcBreakers.size());
193 for (
auto *user : op->getUsers()) {
194 while (user->getParentOp() != op->getParentOp())
195 user = user->getParentOp();
196 auto it = faninMasks.find(user);
197 if (it != faninMasks.end())
201 auto duplicateOp = faninMasks.insert({op,
mask});
203 assert(duplicateOp.second &&
"duplicate op in order");
207 faninMaskGroups.clear();
208 for (
auto [op, mask] : faninMasks)
210 faninMaskGroups[
mask].insert(op);
211 LLVM_DEBUG(llvm::dbgs() <<
"- Found " << faninMaskGroups.size()
212 <<
" fanin mask groups\n");
217void Converter::extractArcs(
HWModuleOp module) {
218 DenseMap<Value, Value> valueMapping;
219 SmallVector<Value> inputs;
220 SmallVector<Value> outputs;
221 SmallVector<Type> inputTypes;
222 SmallVector<Type> outputTypes;
223 SmallVector<std::pair<OpOperand *, unsigned>> externalUses;
226 for (
auto &group : faninMaskGroups) {
227 auto &opSet = group.second;
228 OpBuilder builder(module);
230 auto block = std::make_unique<Block>();
231 builder.setInsertionPointToStart(block.get());
232 valueMapping.clear();
237 externalUses.clear();
239 Operation *lastOp =
nullptr;
241 for (
auto *op : postOrder) {
242 if (!opSet.contains(op))
247 for (
auto &operand : op->getOpOperands()) {
248 if (opSet.contains(operand.get().getDefiningOp()))
250 auto &mapped = valueMapping[operand.get()];
252 mapped = block->addArgument(operand.get().getType(),
253 operand.get().getLoc());
254 inputs.push_back(operand.get());
255 inputTypes.push_back(mapped.getType());
259 for (
auto result : op->getResults()) {
260 bool anyExternal =
false;
261 for (
auto &use : result.getUses()) {
262 if (!opSet.contains(use.getOwner())) {
264 externalUses.push_back({&use, outputs.size()});
268 outputs.push_back(result);
269 outputTypes.push_back(result.getType());
274 arc::OutputOp::create(builder, lastOp->getLoc(), outputs);
277 builder.setInsertionPoint(module);
279 DefineOp::create(builder, lastOp->getLoc(),
280 builder.getStringAttr(globalNamespace.newName(
281 module.getModuleName() +
"_arc")),
282 builder.getFunctionType(inputTypes, outputTypes));
283 defOp.getBody().push_back(block.release());
287 builder.setInsertionPoint(module.getBodyBlock()->getTerminator());
288 auto arcOp = CallOp::create(builder, lastOp->getLoc(), defOp, inputs);
289 arcUses.push_back(arcOp);
290 for (
auto [use, resultIdx] : externalUses)
291 use->set(arcOp.getResult(resultIdx));
295LogicalResult Converter::absorbRegs(
HWModuleOp module) {
299 unsigned numTrivialRegs = 0;
300 for (
auto callOp : arcUses) {
301 auto stateOp = dyn_cast<StateOp>(callOp.getOperation());
302 Value clock = stateOp ? stateOp.getClock() : Value{};
304 SmallVector<Value> initialValues;
305 SmallVector<seq::CompRegOp> absorbedRegs;
306 SmallVector<Attribute> absorbedNames(callOp->getNumResults(), {});
307 if (
auto names = callOp->getAttrOfType<ArrayAttr>(
"names"))
308 absorbedNames.assign(names.getValue().begin(), names.getValue().end());
313 bool isTrivial =
true;
314 for (
auto result : callOp->getResults()) {
315 if (!result.hasOneUse()) {
319 auto regOp = dyn_cast<seq::CompRegOp>(result.use_begin()->getOwner());
320 if (!regOp || regOp.getInput() != result ||
321 (clock && clock != regOp.getClk())) {
326 clock = regOp.getClk();
327 reset = regOp.getReset();
331 Value resetValue = regOp.getResetValue();
332 Operation *op = resetValue.getDefiningOp();
334 return regOp->emitOpError(
335 "is reset by an input; not supported by ConvertToArcs");
336 if (
auto constant = dyn_cast<hw::ConstantOp>(op)) {
337 if (constant.getValue() != 0)
338 return regOp->emitOpError(
"is reset to a constant non-zero value; "
339 "not supported by ConvertToArcs");
341 return regOp->emitOpError(
"is reset to a value that is not clearly "
342 "constant; not supported by ConvertToArcs");
349 absorbedRegs.push_back(regOp);
353 absorbedNames[result.getResultNumber()] = regOp.getNameAttr();
358 arcUses[outIdx++] = callOp;
367 auto arc = dyn_cast<StateOp>(callOp.getOperation());
369 arc.getClockMutable().assign(clock);
370 arc.setLatency(
arc.getLatency() + 1);
372 mlir::IRRewriter rewriter(module->getContext());
373 rewriter.setInsertionPoint(callOp);
374 arc = rewriter.replaceOpWithNewOp<StateOp>(
375 callOp.getOperation(),
376 llvm::cast<SymbolRefAttr>(callOp.getCallableForCallee()),
377 callOp->getResultTypes(), clock, Value{}, 1, callOp.getArgOperands());
382 return arc.emitError(
383 "StateOp tried to infer reset from CompReg, but already "
385 arc.getResetMutable().assign(reset);
388 bool onlyDefaultInitializers =
389 llvm::all_of(initialValues, [](
auto val) ->
bool {
return !val; });
391 if (!onlyDefaultInitializers) {
392 if (!
arc.getInitials().empty()) {
393 return arc.emitError(
394 "StateOp tried to infer initial values from CompReg, but already "
395 "had an initial value.");
398 for (
unsigned i = 0; i < initialValues.size(); ++i) {
399 if (!initialValues[i]) {
400 OpBuilder zeroBuilder(
arc);
403 zeroBuilder.getIntegerAttr(
arc.getResult(i).getType(), 0));
406 arc.getInitialsMutable().assign(initialValues);
409 if (tapRegisters && llvm::any_of(absorbedNames, [](
auto name) {
410 return !cast<StringAttr>(name).getValue().empty();
412 arc->setAttr(
"names", ArrayAttr::get(module.getContext(), absorbedNames));
413 for (
auto [arcResult, reg] :
llvm::zip(
arc.getResults(), absorbedRegs)) {
414 auto it = arcBreakerIndices.find(reg);
415 arcBreakers[it->second] = {};
416 arcBreakerIndices.erase(it);
417 reg.replaceAllUsesWith(arcResult);
421 if (numTrivialRegs > 0)
422 LLVM_DEBUG(llvm::dbgs() <<
"- Trivially converted " << numTrivialRegs
423 <<
" regs to arcs\n");
424 arcUses.truncate(outIdx);
431 for (
auto *op : arcBreakers)
432 if (auto regOp = dyn_cast_or_null<
seq::CompRegOp>(op)) {
433 regsByInput[{regOp.getClk(), regOp.getReset(),
434 regOp.getInput().getDefiningOp()}]
438 unsigned numMappedRegs = 0;
439 for (
auto [clockAndResetAndOp, regOps] : regsByInput) {
440 numMappedRegs += regOps.size();
441 OpBuilder builder(module);
442 auto block = std::make_unique<Block>();
443 builder.setInsertionPointToStart(block.get());
445 SmallVector<Value> inputs;
446 SmallVector<Value> outputs;
447 SmallVector<Attribute> names;
448 SmallVector<Type> types;
449 SmallVector<Value> initialValues;
451 SmallVector<unsigned> regToOutputMapping;
452 for (
auto regOp : regOps) {
453 auto it = mapping.find(regOp.getInput());
454 if (it == mapping.end()) {
455 it = mapping.insert({regOp.getInput(), inputs.size()}).first;
456 inputs.push_back(regOp.getInput());
457 types.push_back(regOp.getType());
458 outputs.push_back(block->addArgument(regOp.getType(), regOp.getLoc()));
459 names.push_back(regOp->getAttrOfType<StringAttr>(
"name"));
463 regToOutputMapping.push_back(it->second);
466 auto loc = regOps.back().getLoc();
467 arc::OutputOp::create(builder, loc, outputs);
469 builder.setInsertionPoint(module);
470 auto defOp = DefineOp::create(builder, loc,
471 builder.getStringAttr(globalNamespace.newName(
472 module.getModuleName() +
"_arc")),
473 builder.getFunctionType(types, types));
474 defOp.getBody().push_back(block.release());
476 builder.setInsertionPoint(module.getBodyBlock()->getTerminator());
478 bool onlyDefaultInitializers =
479 llvm::all_of(initialValues, [](
auto val) ->
bool {
return !val; });
481 if (onlyDefaultInitializers)
482 initialValues.clear();
484 for (
unsigned i = 0; i < initialValues.size(); ++i) {
485 if (!initialValues[i])
487 loc, builder.getIntegerAttr(types[i], 0));
491 StateOp::create(builder, loc, defOp, std::get<0>(clockAndResetAndOp),
492 Value{}, 1, inputs, initialValues);
493 auto reset = std::get<1>(clockAndResetAndOp);
495 arcOp.getResetMutable().assign(reset);
496 if (tapRegisters && llvm::any_of(names, [](
auto name) {
497 return !cast<StringAttr>(name).getValue().empty();
499 arcOp->setAttr(
"names", builder.getArrayAttr(names));
500 for (
auto [reg, resultIdx] :
llvm::zip(regOps, regToOutputMapping)) {
501 reg.replaceAllUsesWith(arcOp.getResult(resultIdx));
506 if (numMappedRegs > 0)
507 LLVM_DEBUG(llvm::dbgs() <<
"- Mapped " << numMappedRegs <<
" regs to "
508 << regsByInput.size() <<
" shuffling arcs\n");
518static LogicalResult
convert(llhd::CombinationalOp op,
519 llhd::CombinationalOp::Adaptor adaptor,
520 ConversionPatternRewriter &rewriter,
521 const TypeConverter &converter) {
523 SmallVector<Type> resultTypes;
524 if (failed(converter.convertTypes(op.getResultTypes(), resultTypes)))
528 auto cloneIntoBody = [](Operation *op) {
529 return op->hasTrait<OpTrait::ConstantLike>();
532 mlir::makeRegionIsolatedFromAbove(rewriter, op.getBody(), cloneIntoBody);
536 ExecuteOp::create(rewriter, op.getLoc(), resultTypes, operands);
537 executeOp.getBody().takeBody(op.getBody());
538 rewriter.replaceOp(op, executeOp.getResults());
543static LogicalResult
convert(llhd::YieldOp op, llhd::YieldOp::Adaptor adaptor,
544 ConversionPatternRewriter &rewriter) {
545 rewriter.replaceOpWithNewOp<arc::OutputOp>(op, adaptor.getOperands());
554#define GEN_PASS_DEF_CONVERTTOARCSPASS
555#include "circt/Conversion/Passes.h.inc"
559struct ConvertToArcsPass
560 :
public circt::impl::ConvertToArcsPassBase<ConvertToArcsPass> {
561 using ConvertToArcsPassBase::ConvertToArcsPassBase;
562 void runOnOperation()
override;
566void ConvertToArcsPass::runOnOperation() {
568 TypeConverter converter;
569 converter.addConversion([](Type type) {
return type; });
579 ConversionTarget target(getContext());
580 target.addIllegalOp<llhd::CombinationalOp, llhd::YieldOp>();
581 target.markUnknownOpDynamicallyLegal([](Operation *) {
return true; });
584 ConversionConfig config;
585 config.allowPatternRollback =
false;
588 if (failed(applyPartialConversion(getOperation(), target, std::move(
patterns),
590 emitError(getOperation().
getLoc()) <<
"conversion to arcs failed";
591 return signalPassFailure();
596 outliner.tapRegisters = tapRegisters;
597 if (failed(outliner.run(getOperation())))
598 return signalPassFailure();
assert(baseType &&"element must be base type")
static LogicalResult convertInitialValue(seq::CompRegOp reg, SmallVectorImpl< Value > &values)
static LogicalResult convert(llhd::CombinationalOp op, llhd::CombinationalOp::Adaptor adaptor, ConversionPatternRewriter &rewriter, const TypeConverter &converter)
llhd.combinational -> arc.execute
static bool isArcBreakingOp(Operation *op)
static Location getLoc(DefSlot slot)
static Block * getBodyBlock(FModuleLike mod)
Extension of RewritePatternSet that allows adding matchAndRewrite functions with op adaptors and Conv...
A namespace that is used to store existing names and generate new names in some scope within the IR.
StringAttr getName(ArrayAttr names, size_t idx)
Return the name at the specified index of the ArrayAttr or null if it cannot be determined.
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.
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)