14#include "mlir/IR/Matchers.h"
15#include "llvm/Support/Debug.h"
17#define DEBUG_TYPE "llhd-combine-drives"
21#define GEN_PASS_DEF_COMBINEDRIVESPASS
22#include "circt/Dialect/LLHD/LLHDPasses.h.inc"
32using llvm::SpecificBumpPtrAllocator;
38 return TypeSwitch<Type, unsigned>(cast<RefType>(type).getNestedType())
39 .Case<IntegerType>([](
auto type) {
return type.getWidth(); })
40 .Case<hw::ArrayType>([](
auto type) {
return type.getNumElements(); })
41 .Case<hw::StructType>([](
auto type) {
return type.getElements().size(); })
42 .Case<hw::UnionType>([](
auto type) {
return type.getElements().size(); })
43 .Default([](
auto) {
return 0; });
88 Signal *signal =
nullptr;
92 explicit operator bool()
const {
return signal !=
nullptr; }
111 Signal *parent =
nullptr;
113 unsigned indexInParent = 0;
116 SmallVector<Signal *> subsignals;
121 SmallVector<ValueSlice> slices;
124 SmallVector<DriveOp, 2> completeDrives;
127 explicit Signal(Value root) : value(root) {}
129 Signal(Value value, Signal *parent,
unsigned indexInParent)
130 : value(value), parent(parent), indexInParent(indexInParent) {}
135struct ModuleContext {
136 ModuleContext(
HWModuleOp moduleOp) : moduleOp(moduleOp) {}
140 SignalSlice traceProjection(Value value);
141 SignalSlice traceProjectionImpl(Value value);
142 Signal *internSignal(Value root);
143 Signal *internSignal(Value value, Signal *parent,
unsigned index);
146 void aggregateDrives(Signal &signal);
147 void addDefaultDriveSlices(Signal &signal,
148 SmallVectorImpl<DriveSlice> &slices);
149 void aggregateDriveSlices(Signal &signal, Value driveDelay, Value driveEnable,
150 ArrayRef<DriveSlice> slices);
155 DenseMap<Value, SignalSlice> projections;
157 SmallVector<Signal *> rootSignals;
162 using SignalKey = std::pair<PointerUnion<Value, Signal *>,
unsigned>;
163 SpecificBumpPtrAllocator<Signal> signalAlloc;
164 DenseMap<SignalKey, Signal *> internedSignals;
170 const Signal &signal) {
172 return os << *signal.parent <<
"[" << signal.indexInParent <<
"]";
173 signal.value.printAsOperand(os, OpPrintingFlags().useLocalScope());
178static llvm::raw_ostream &
operator<<(llvm::raw_ostream &os, SignalSlice slice) {
180 return os <<
"<null-slice>";
181 return os << *slice.signal <<
"[" << slice.offset <<
".."
182 << (slice.offset + slice.length) <<
"]";
192SignalSlice ModuleContext::traceProjection(Value value) {
194 if (
auto it = projections.find(value); it != projections.end())
198 auto projection = traceProjectionImpl(value);
200 projection.signal->slices.push_back(
201 ValueSlice{value, projection.offset, projection.length});
202 projections.insert({value, projection});
203 LLVM_DEBUG(llvm::dbgs() <<
"- Traced " << value <<
" to " << projection
209SignalSlice ModuleContext::traceProjectionImpl(Value value) {
214 if (
auto op = value.getDefiningOp<SigExtractOp>()) {
215 auto slice = traceProjection(op.getInput());
218 IntegerAttr offsetAttr;
219 if (!matchPattern(op.getLowBit(), m_Constant(&offsetAttr)))
221 slice.offset += offsetAttr.getValue().getZExtValue();
222 slice.length =
getLength(value.getType());
226 if (
auto op = value.getDefiningOp<SigArraySliceOp>()) {
227 auto slice = traceProjection(op.getInput());
230 IntegerAttr offsetAttr;
231 if (!matchPattern(op.getLowIndex(), m_Constant(&offsetAttr)))
233 slice.offset += offsetAttr.getValue().getZExtValue();
234 slice.length =
getLength(value.getType());
241 if (
auto op = value.getDefiningOp<SigArrayGetOp>()) {
242 auto input = traceProjection(op.getInput());
245 IntegerAttr indexAttr;
246 if (!matchPattern(op.getIndex(), m_Constant(&indexAttr)))
248 unsigned offset = input.offset + indexAttr.getValue().getZExtValue();
250 slice.signal = internSignal(value, input.signal, offset);
251 slice.length =
getLength(value.getType());
255 if (
auto op = value.getDefiningOp<SigStructExtractOp>()) {
256 auto input = traceProjection(op.getInput());
259 auto type = cast<RefType>(op.getInput().getType()).getNestedType();
260 if (
auto structType = hw::type_dyn_cast<hw::StructType>(type)) {
261 assert(input.offset == 0);
262 assert(input.length == structType.getElements().size());
263 unsigned index = *structType.getFieldIndex(op.getFieldAttr());
265 slice.signal = internSignal(value, input.signal, index);
266 slice.length =
getLength(value.getType());
269 auto unionType = hw::type_cast<hw::UnionType>(type);
270 assert(input.offset == 0);
271 assert(input.length == unionType.getElements().size());
272 unsigned index = *unionType.getFieldIndex(op.getFieldAttr());
274 slice.signal = internSignal(value, input.signal, index);
275 slice.length =
getLength(value.getType());
282 slice.signal = internSignal(value);
283 slice.length =
getLength(value.getType());
290Signal *ModuleContext::internSignal(Value root) {
291 auto &slot = internedSignals[{root, 0}];
293 slot =
new (signalAlloc.Allocate()) Signal(root);
294 rootSignals.push_back(slot);
302Signal *ModuleContext::internSignal(Value value, Signal *parent,
304 auto &slot = internedSignals[{parent, index}];
306 slot =
new (signalAlloc.Allocate()) Signal(value, parent, index);
307 parent->subsignals.push_back(slot);
321void ModuleContext::aggregateDrives(Signal &signal) {
328 SmallPtrSet<Operation *, 8> knownDrives;
329 auto addDrive = [&](DriveOp op,
unsigned offset,
unsigned length) {
330 knownDrives.insert(op);
331 drives[{op.getTime(), op.getEnable()}].push_back(
332 DriveSlice{op, op.getValue(), offset, length});
334 for (
auto *subsignal : signal.subsignals) {
335 aggregateDrives(*subsignal);
343 for (
auto driveOp : subsignal->completeDrives)
344 addDrive(driveOp, subsignal->indexInParent, 0);
348 for (
auto slice : signal.slices) {
349 for (
auto &use : slice.value.getUses()) {
350 auto driveOp = dyn_cast<DriveOp>(use.getOwner());
351 if (driveOp && use.getOperandNumber() == 0 &&
352 driveOp->getBlock() == slice.value.getParentBlock())
353 addDrive(driveOp, slice.offset, slice.length);
361 worklist.insert(signal.value);
362 bool hasUnknownUses =
false;
363 while (!worklist.empty() && !hasUnknownUses) {
364 auto value = worklist.pop_back_val();
365 for (
auto *user : value.getUsers()) {
366 if (isa<ProbeOp>(user))
368 if (isa<DriveOp>(user) && knownDrives.contains(user))
370 if (isa<SigExtractOp, SigStructExtractOp, SigArrayGetOp, SigArraySliceOp>(
372 worklist.insert(user->getResult(0));
375 hasUnknownUses =
true;
383 if (!hasUnknownUses && drives.size() == 1) {
384 auto &slices = drives.begin()->second;
385 addDefaultDriveSlices(signal, slices);
390 for (
auto &[key, slices] : drives) {
391 llvm::sort(slices, [](
auto &a,
auto &b) {
return a.offset < b.offset; });
392 aggregateDriveSlices(signal, key.first, key.second, slices);
399void ModuleContext::addDefaultDriveSlices(Signal &signal,
400 SmallVectorImpl<DriveSlice> &slices) {
401 auto type = cast<RefType>(signal.value.getType()).getNestedType();
404 llvm::sort(slices, [](
auto &a,
auto &b) {
return a.offset < b.offset; });
409 bool anyOverlaps =
false;
410 bool needSeparateFields = isa<hw::StructType>(type);
411 SmallVector<DriveSlice> gapSlices;
412 auto fillGap = [&](
unsigned from,
unsigned to) {
419 if (needSeparateFields) {
420 for (
auto idx = from; idx < to; ++idx)
421 gapSlices.push_back(DriveSlice{DriveOp{}, Value{}, idx, 0});
423 gapSlices.push_back(DriveSlice{DriveOp{}, Value{}, from, to - from});
431 if (hw::type_isa<hw::UnionType>(type)) {
433 gapSlices.push_back(DriveSlice{DriveOp{}, Value{}, 0, 0});
435 unsigned expectedOffset = 0;
436 for (
auto slice : slices) {
437 fillGap(expectedOffset, slice.offset);
438 expectedOffset = slice.offset + std::max<unsigned>(1, slice.length);
442 fillGap(expectedOffset,
getLength(signal.value.getType()));
447 if (anyOverlaps || gapSlices.empty())
462 auto signalOp = signal.value.getDefiningOp<SignalOp>();
465 auto defaultValue = signalOp.getInit();
468 ImplicitLocOpBuilder builder(signal.value.getLoc(),
469 signal.value.getContext());
470 builder.setInsertionPointAfterValue(signal.value);
472 for (
auto &slice : gapSlices) {
473 LLVM_DEBUG(llvm::dbgs()
474 <<
"- Filling gap " << signal <<
"[" << slice.offset <<
".."
475 << (slice.offset + slice.length) <<
"] with initial value\n");
478 if (
auto intType = dyn_cast<IntegerType>(type)) {
482 defaultValue, slice.offset);
487 if (
auto structType = hw::type_dyn_cast<hw::StructType>(type)) {
488 assert(slice.length == 0);
490 builder, defaultValue, structType.getElements()[slice.offset]);
495 if (
auto unionType = hw::type_dyn_cast<hw::UnionType>(type)) {
496 assert(slice.offset == 0 && slice.length == 0);
497 slice.value = hw::UnionExtractOp::create(builder, defaultValue, 0);
502 if (
auto arrayType = dyn_cast<hw::ArrayType>(type)) {
506 APInt(llvm::Log2_64_Ceil(arrayType.getNumElements()), slice.offset));
508 builder, hw::ArrayType::get(arrayType.getElementType(), slice.length),
509 defaultValue, offset);
517 slices.append(gapSlices.begin(), gapSlices.end());
522void ModuleContext::aggregateDriveSlices(Signal &signal, Value driveDelay,
524 ArrayRef<DriveSlice> slices) {
525 auto type = cast<RefType>(signal.value.getType()).getNestedType();
529 if (hw::type_isa<hw::UnionType>(type)) {
530 if (slices.size() != 1) {
531 LLVM_DEBUG(llvm::dbgs()
532 <<
"- Union " << signal <<
" not uniquely driven\n");
536 unsigned expectedOffset = 0;
537 for (
auto slice : slices) {
538 assert(slice.value &&
"all slices must have an assigned value");
539 if (slice.offset != expectedOffset) {
547 expectedOffset += std::max<unsigned>(1, slice.length);
549 if (expectedOffset !=
getLength(signal.value.getType())) {
550 LLVM_DEBUG(llvm::dbgs()
551 <<
"- Signal " << signal <<
" not completely driven\n");
559 if (slices.size() == 1 && slices[0].length != 0 && slices[0].op &&
560 !hw::type_isa<hw::UnionType>(type)) {
561 signal.completeDrives.push_back(slices[0].op);
565 llvm::dbgs() <<
"- Aggregating " << signal <<
" drives (delay ";
566 driveDelay.printAsOperand(llvm::dbgs(), OpPrintingFlags().useLocalScope());
568 llvm::dbgs() <<
" if ";
569 driveEnable.printAsOperand(llvm::dbgs(),
570 OpPrintingFlags().useLocalScope());
572 llvm::dbgs() <<
")\n";
576 ImplicitLocOpBuilder builder(signal.value.getLoc(),
577 signal.value.getContext());
578 builder.setInsertionPointAfterValue(signal.value);
581 if (
auto intType = dyn_cast<IntegerType>(type)) {
585 SmallVector<Value> operands;
586 for (
auto slice : slices)
587 operands.push_back(slice.value);
588 std::reverse(operands.begin(), operands.end());
589 result = comb::ConcatOp::create(builder, operands);
590 LLVM_DEBUG(llvm::dbgs() <<
" - Created " << result <<
"\n");
594 if (
auto structType = hw::type_dyn_cast<hw::StructType>(type)) {
597 SmallVector<Value> operands;
598 for (
auto slice : slices)
599 operands.push_back(slice.value);
601 LLVM_DEBUG(llvm::dbgs() <<
" - Created " << result <<
"\n");
605 if (
auto unionType = hw::type_dyn_cast<hw::UnionType>(type)) {
608 assert(slices.size() == 1);
609 result = hw::UnionCreateOp::create(builder, unionType, slices[0].offset,
611 LLVM_DEBUG(llvm::dbgs() <<
" - Created " << result <<
"\n");
615 if (
auto arrayType = dyn_cast<hw::ArrayType>(type)) {
619 SmallVector<Value> scalars;
620 SmallVector<Value> aggregates;
621 auto flushScalars = [&] {
624 std::reverse(scalars.begin(), scalars.end());
626 aggregates.push_back(aggregate);
628 LLVM_DEBUG(llvm::dbgs() <<
" - Created " << aggregate <<
"\n");
630 for (
auto slice : slices) {
631 if (slice.length == 0) {
632 scalars.push_back(slice.value);
635 aggregates.push_back(slice.value);
642 result = aggregates.back();
643 if (aggregates.size() != 1) {
644 std::reverse(aggregates.begin(), aggregates.end());
646 LLVM_DEBUG(llvm::dbgs() <<
" - Created " << result <<
"\n");
653 DriveOp::create(builder, signal.value, result, driveDelay, driveEnable);
654 signal.completeDrives.push_back(driveOp);
655 LLVM_DEBUG(llvm::dbgs() <<
" - Created " << driveOp <<
"\n");
658 for (
auto slice : slices) {
661 LLVM_DEBUG(llvm::dbgs() <<
" - Removed " << slice.op <<
"\n");
662 pruner.eraseNow(slice.op);
671struct CombineDrivesPass
672 :
public llhd::impl::CombineDrivesPassBase<CombineDrivesPass> {
673 void runOnOperation()
override;
677void CombineDrivesPass::runOnOperation() {
678 LLVM_DEBUG(llvm::dbgs() <<
"Combining drives in "
679 << getOperation().getModuleNameAttr() <<
"\n");
680 ModuleContext
context(getOperation());
684 if (isa<SigExtractOp, SigArraySliceOp, SigArrayGetOp, SigStructExtractOp>(
686 context.traceProjection(op.getResult(0));
689 for (
auto *signal :
context.rootSignals)
690 context.aggregateDrives(*signal);
assert(baseType &&"element must be base type")
static unsigned getLength(Type type)
Determine the number of elements in a type.
static std::unique_ptr< Context > context
static Block * getBodyBlock(FModuleLike mod)
create(elements, Type result_type=None)
create(array_value, low_index, ret_type)
create(elements, Type result_type=None)
OS & operator<<(OS &os, const InnerSymTarget &target)
Printing InnerSymTarget's.
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.
Utility that tracks operations that have potentially become unused and allows them to be cleaned up a...