26#include "mlir/IR/IRMapping.h"
27#include "mlir/Transforms/DialectConversion.h"
31#define GEN_PASS_DEF_SIMPLIFYREFS
32#include "circt/Dialect/Moore/MoorePasses.h.inc"
43static void collectOperands(Value operand, SmallVectorImpl<Value> &operands,
44 ConversionPatternRewriter &rewriter) {
45 if (
auto concatRefOp = operand.getDefiningOp<ConcatRefOp>()) {
47 if (std::distance(concatRefOp->getUsers().begin(),
48 concatRefOp->getUsers().end()) == 1) {
49 rewriter.eraseOp(concatRefOp);
51 for (
auto nestedOperand : concatRefOp.getValues())
52 collectOperands(nestedOperand, operands, rewriter);
55 operands.push_back(operand);
58template <
typename OpTy>
61 using OpAdaptor =
typename OpTy::Adaptor;
64 matchAndRewrite(OpTy op, OpAdaptor adaptor,
65 ConversionPatternRewriter &rewriter)
const override {
67 SmallVector<Value, 4> operands;
68 collectOperands(op.getDst(), operands, rewriter);
70 cast<UnpackedType>(op.getSrc().getType()).getBitSize().value();
74 for (
auto operand : operands) {
75 auto type = cast<RefType>(operand.getType()).getNestedType();
76 auto width = type.getBitSize().value();
78 rewriter.setInsertionPoint(op);
83 auto extract = ExtractOp::create(rewriter, op.getLoc(), type, op.getSrc(),
88 srcWidth = srcWidth - width;
94 mapping.map(op.getDst(), operand);
95 mapping.map(op.getSrc(), Value(extract));
96 rewriter.clone(*op.getOperation(), mapping);
104 Value field =
nullptr;
111 SmallVector<FieldInfo> &fields,
112 ConversionPatternRewriter &rewriter,
113 uint32_t initialOffset = 0) {
115 cast<StructType>(cast<RefType>(structRef.getType()).getNestedType());
116 uint32_t offset = initialOffset;
118 for (
auto &member :
llvm::reverse(structType.getMembers())) {
119 auto fieldRef = StructExtractRefOp::create(rewriter, structRef.getLoc(),
120 RefType::get(member.type),
121 member.name, structRef);
123 auto fieldSize = member.type.getBitSize();
125 return mlir::emitError(structRef.getLoc())
126 <<
"unsupported: field with unknown size in struct flattening";
128 if (isa<StructType>(member.type)) {
129 auto result =
collectFields(fieldRef, fields, rewriter, offset);
132 }
else if (isa<UnionType>(member.type)) {
133 return mlir::emitError(structRef.getLoc())
134 <<
"unsupported: union member in struct flattening";
136 fields.push_back({fieldRef, *fieldSize, offset});
138 offset += *fieldSize;
143template <
typename OpTy>
146 using OpAdaptor =
typename OpTy::Adaptor;
149 matchAndRewrite(OpTy op, OpAdaptor adaptor,
150 ConversionPatternRewriter &rewriter)
const override {
151 Value dst = op.getDst();
155 auto extractRefOp = dyn_cast<ExtractRefOp>(dst.getDefiningOp());
156 auto baseStructType = dyn_cast_if_present<StructType>(
157 cast<RefType>(extractRefOp.getInput().getType()).getNestedType());
162 uint32_t extractedLow = extractRefOp.getLowBit();
164 cast<RefType>(dst.getType()).getNestedType().getBitSize();
166 return mlir::emitError(op.getLoc())
167 <<
"unsupported: found field with unknown size in struct "
170 rewriter.setInsertionPoint(op);
171 Location loc = op.getLoc();
174 SmallVector<FieldInfo> fields;
175 if (failed(
collectFields(extractRefOp.getInput(), fields, rewriter))) {
180 SmallVector<Value> relevantFields;
181 uint32_t offsetInStruct = extractedLow;
182 uint32_t remaining = *targetWidth;
183 for (
auto &field : fields) {
187 if (field.offset <= offsetInStruct &&
188 field.offset + field.size > offsetInStruct) {
189 uint32_t offsetInField = offsetInStruct - field.offset;
190 uint32_t remainingInField = field.size - offsetInField;
191 uint32_t sizeToExtract = std::min(remaining, remainingInField);
193 auto fieldType = cast<RefType>(field.field.getType()).getNestedType();
195 IntType::get(rewriter.getContext(), sizeToExtract,
196 cast<PackedType>(fieldType).getDomain());
198 ExtractRefOp::create(rewriter, loc, RefType::get(extractType),
199 field.field, offsetInField);
200 relevantFields.push_back(newLhs);
202 offsetInStruct += sizeToExtract;
203 remaining -= sizeToExtract;
207 if (relevantFields.size() == 0)
208 return mlir::emitError(op.getLoc()) <<
"Struct extract is out of range";
211 std::reverse(relevantFields.begin(), relevantFields.end());
214 relevantFields.size() == 1
215 ? relevantFields.front()
216 : Value(ConcatRefOp::create(rewriter, loc, relevantFields));
219 mapping.map(op.getDst(), finalRef);
220 rewriter.clone(*op.getOperation(), mapping);
222 rewriter.eraseOp(op);
223 rewriter.eraseOp(extractRefOp);
232 matchAndRewrite(DynQueueRefElementOp op, OpAdaptor adaptor,
233 ConversionPatternRewriter &rewriter)
const override {
236 for (
auto *consumer : op->getUsers()) {
237 if (isa<BlockingAssignOp>(consumer)) {
239 auto assignOp = cast<BlockingAssignOp>(consumer);
241 rewriter.setInsertionPoint(consumer);
242 moore::QueueSetOp::create(rewriter, op->getLoc(), op.getInput(),
243 op.getIndex(), assignOp.getSrc());
245 rewriter.eraseOp(assignOp);
247 return mlir::emitError(op.getLoc())
248 <<
"Queue element reference couldn't be reduced to setting the "
249 "value at an index: consuming op "
250 << consumer <<
" is not supported";
254 rewriter.eraseOp(op);
260struct AssocArrayRefLowering
265 matchAndRewrite(AssocArrayExtractRefOp op, OpAdaptor adaptor,
266 ConversionPatternRewriter &rewriter)
const override {
267 for (
auto *consumer : op->getUsers()) {
268 if (isa<BlockingAssignOp>(consumer)) {
269 auto assignOp = cast<BlockingAssignOp>(consumer);
270 rewriter.setInsertionPoint(consumer);
271 moore::AssocArraySetOp::create(rewriter, op->getLoc(), op.getInput(),
272 op.getIndex(), assignOp.getSrc());
273 rewriter.eraseOp(assignOp);
275 return mlir::emitError(op.getLoc())
276 <<
"Associative array element reference couldn't be reduced "
277 "to setting the value at an index: consuming op "
278 << consumer <<
" is not supported";
282 rewriter.eraseOp(op);
287struct SimplifyRefsPass
288 :
public circt::moore::impl::SimplifyRefsBase<SimplifyRefsPass> {
289 void runOnOperation()
override;
295 return std::make_unique<SimplifyRefsPass>();
298void SimplifyRefsPass::runOnOperation() {
299 MLIRContext &
context = getContext();
300 ConversionTarget target(
context);
302 target.addDynamicallyLegalOp<ContinuousAssignOp, BlockingAssignOp,
303 NonBlockingAssignOp, DelayedContinuousAssignOp,
304 DelayedNonBlockingAssignOp>([](
auto op) {
306 op->getOperand(0).template getDefiningOp<ExtractRefOp>();
310 cast<RefType>(extractRefOp.getInput().getType()).getNestedType();
311 return !isa<StructType>(nestedType);
314 target.markUnknownOpDynamicallyLegal([](Operation *) {
return true; });
316 RewritePatternSet extractRefOnStructPatterns(&
context);
317 extractRefOnStructPatterns
318 .add<StructExtractLowering<ContinuousAssignOp>,
319 StructExtractLowering<BlockingAssignOp>,
320 StructExtractLowering<NonBlockingAssignOp>,
321 StructExtractLowering<DelayedContinuousAssignOp>,
322 StructExtractLowering<DelayedNonBlockingAssignOp>>(&
context);
324 if (failed(applyFullConversion(getOperation(), target,
325 std::move(extractRefOnStructPatterns)))) {
330 target.addDynamicallyLegalOp<ContinuousAssignOp, BlockingAssignOp,
331 NonBlockingAssignOp, DelayedContinuousAssignOp,
332 DelayedNonBlockingAssignOp>([](
auto op) {
333 return !op->getOperand(0).template getDefiningOp<ConcatRefOp>();
336 RewritePatternSet concatRefPatterns(&
context);
337 concatRefPatterns.add<ConcatRefLowering<ContinuousAssignOp>,
338 ConcatRefLowering<BlockingAssignOp>,
339 ConcatRefLowering<NonBlockingAssignOp>,
340 ConcatRefLowering<DelayedContinuousAssignOp>,
341 ConcatRefLowering<DelayedNonBlockingAssignOp>>(
344 if (failed(applyFullConversion(getOperation(), target,
345 std::move(concatRefPatterns)))) {
352 RewritePatternSet queueRefPatterns(&
context);
353 target.addIllegalOp<DynQueueRefElementOp>();
354 queueRefPatterns.add<QueueRefLowering>(&
context);
355 if (failed(applyPartialConversion(getOperation(), target,
356 std::move(queueRefPatterns))))
361 RewritePatternSet assocArrayRefPatterns(&
context);
362 target.addIllegalOp<AssocArrayExtractRefOp>();
363 assocArrayRefPatterns.add<AssocArrayRefLowering>(&
context);
364 if (failed(applyPartialConversion(getOperation(), target,
365 std::move(assocArrayRefPatterns))))
static std::unique_ptr< Context > context
static Attribute collectFields(MLIRContext *context, ArrayRef< Attribute > operands)
std::unique_ptr< mlir::Pass > createSimplifyRefsPass()
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.