13#include "mlir/Pass/Pass.h"
14#include "mlir/Transforms/DialectConversion.h"
15#include "llvm/ADT/APInt.h"
19#define GEN_PASS_DEF_HWAGGREGATETOCOMB
20#include "circt/Dialect/HW/Passes.h.inc"
30template <
typename OpTy>
35 matchAndRewrite(OpTy op, OpAdaptor adaptor,
36 ConversionPatternRewriter &rewriter)
const override {
37 rewriter.replaceOpWithNewOp<
comb::ConcatOp>(op, adaptor.getInputs());
42struct HWUnionCreateOpConversion
47 matchAndRewrite(hw::UnionCreateOp op, OpAdaptor adaptor,
48 ConversionPatternRewriter &rewriter)
const override {
49 hw::UnionType unionTy = op.getType();
51 dyn_cast_or_null<IntegerType>(typeConverter->convertType(unionTy));
53 return rewriter.notifyMatchFailure(op.getLoc(),
54 "Failed to convert union to integer");
56 auto inputBitWidth = hw::getBitWidth(adaptor.getInput().getType());
57 if (inputBitWidth < 0)
58 return rewriter.notifyMatchFailure(op.getLoc(),
59 "Failed to convert input to integer");
62 auto inputIntTy = rewriter.getIntegerType(inputBitWidth);
64 op.getLoc(), inputIntTy, adaptor.getInput());
69 int64_t bitOffset = unionTy.getElements()[op.getFieldIndex()].offset;
70 int64_t prePadding = outputTy.getWidth() - inputBitWidth - bitOffset;
72 auto createZeroCst = [&](Location loc, int64_t bitWidth) -> Value {
74 rewriter.getIntegerType(bitWidth), 0);
77 SmallVector<Value> concatOperands;
79 concatOperands.push_back(createZeroCst(op.getLoc(), prePadding));
80 concatOperands.push_back(inputAsInt);
82 concatOperands.push_back(createZeroCst(op.getLoc(), bitOffset));
85 rewriter.createOrFold<
comb::ConcatOp>(op.getLoc(), concatOperands);
86 rewriter.replaceOp(op, result);
91struct HWUnionExtractOpConversion
96 matchAndRewrite(hw::UnionExtractOp op, OpAdaptor adaptor,
97 ConversionPatternRewriter &rewriter)
const override {
98 hw::UnionType unionTy = op.getInput().getType();
100 auto inputTy = dyn_cast_or_null<IntegerType>(adaptor.getInput().getType());
102 return rewriter.notifyMatchFailure(op.getLoc(),
103 "Failed to convert union to integer");
104 auto outputTy = typeConverter->convertType(op.getType());
106 return rewriter.notifyMatchFailure(
107 op.getLoc(),
"Failed to convert union extract result type");
109 auto resultFieldBits = hw::getBitWidth(outputTy);
110 assert(resultFieldBits >= 0);
111 auto integerValue = adaptor.getInput();
114 if (resultFieldBits < integerValue.getType().getIntOrFloatBitWidth()) {
115 auto bitOffset = unionTy.getElements()[op.getFieldIndex()].offset;
117 rewriter, op->getLoc(), rewriter.getIntegerType(resultFieldBits),
118 integerValue, bitOffset);
123 auto bitcastOp = rewriter.createOrFold<
hw::BitcastOp>(op.getLoc(), outputTy,
126 rewriter.replaceOp(op, bitcastOp);
131struct HWAggregateConstantOpConversion
136 matchAndRewrite(hw::AggregateConstantOp op, OpAdaptor adaptor,
137 ConversionPatternRewriter &rewriter)
const override {
140 if (failed(hw::aggregateAttrToAPInt(op.getType(), adaptor.getFieldsAttr(),
153 ConversionPatternRewriter &rewriter)
const override {
154 SmallVector<Value> results;
155 auto arrayType = cast<hw::ArrayType>(op.getInput().getType());
156 auto elemType = arrayType.getElementType();
158 auto elemWidth = hw::getBitWidth(elemType);
160 return rewriter.notifyMatchFailure(op.getLoc(),
"unknown element width");
162 auto lowered = adaptor.getInput();
163 auto index = adaptor.getIndex();
165 if (matchPattern(index, m_ConstantInt(&constantIndex))) {
166 int64_t maxIndex = std::numeric_limits<int32_t>::max() / elemWidth;
167 if (constantIndex.isSingleWord() &&
168 constantIndex.getZExtValue() <=
static_cast<uint64_t
>(maxIndex)) {
170 op, lowered, constantIndex.getZExtValue() * elemWidth, elemWidth);
177 op.getLoc(), lowered, i * elemWidth, elemWidth));
179 SmallVector<Value> bits;
180 comb::extractBits(rewriter, index, bits);
181 auto result = comb::constructMuxTree(rewriter, op.getLoc(), bits, results,
184 rewriter.replaceOp(op, result);
194 ConversionPatternRewriter &rewriter)
const override {
195 SmallVector<Value> results;
196 auto arrayType = cast<hw::ArrayType>(op.getInput().getType());
197 auto elemType = arrayType.getElementType();
199 auto elemWidth = hw::getBitWidth(elemType);
201 return rewriter.notifyMatchFailure(op.getLoc(),
"unknown element width");
202 auto resultArrayType = cast<hw::ArrayType>(op.getResult().getType());
203 auto resultNumElements = resultArrayType.getNumElements();
205 auto lowered = adaptor.getInput();
206 auto index = adaptor.getLowIndex();
208 if (matchPattern(index, m_ConstantInt(&constantIndex))) {
209 int64_t maxIndex = std::numeric_limits<int32_t>::max() / elemWidth;
210 if (constantIndex.isSingleWord() &&
211 constantIndex.getZExtValue() <=
static_cast<uint64_t
>(maxIndex)) {
213 op, lowered, constantIndex.getZExtValue() * elemWidth,
214 resultNumElements * elemWidth);
219 for (
size_t i = 0; i <=
numElements - resultNumElements; ++i)
221 op.getLoc(), lowered, i * elemWidth, resultNumElements * elemWidth));
223 SmallVector<Value> bits;
224 comb::extractBits(rewriter, index, bits);
225 auto result = comb::constructMuxTree(rewriter, op.getLoc(), bits, results,
228 rewriter.replaceOp(op, result);
237 matchAndRewrite(hw::ArrayInjectOp op, OpAdaptor adaptor,
238 ConversionPatternRewriter &rewriter)
const override {
239 auto arrayType = cast<hw::ArrayType>(op.getInput().getType());
240 auto elemType = arrayType.getElementType();
242 auto elemWidth = hw::getBitWidth(elemType);
244 return rewriter.notifyMatchFailure(op.getLoc(),
"unknown element width");
246 Location loc = op.getLoc();
249 SmallVector<Value> originalElements;
250 auto inputArray = adaptor.getInput();
253 loc, inputArray, i * elemWidth, elemWidth));
258 SmallVector<Value> arrayRows;
260 for (
int injectIdx =
numElements - 1; injectIdx >= 0; --injectIdx) {
261 SmallVector<Value> rowElements;
266 for (
int originalIdx =
numElements - 1; originalIdx >= 0; --originalIdx) {
267 if (originalIdx == injectIdx) {
268 rowElements.push_back(adaptor.getElement());
270 rowElements.push_back(originalElements[originalIdx]);
276 arrayRows.push_back(row);
288 rewriter.replaceOp(op, arrayGetOp);
298 ConversionPatternRewriter &rewriter)
const override {
301 rewriter.replaceOpWithNewOp<
comb::ConcatOp>(op, adaptor.getInput());
311 ConversionPatternRewriter &rewriter)
const override {
312 auto structType = cast<hw::StructType>(op.getInput().getType());
313 auto fieldIndex = op.getFieldIndex();
314 auto elements = structType.getElements();
316 int64_t totalBitWidth = hw::getBitWidth(structType);
317 if (totalBitWidth < 0)
318 return rewriter.notifyMatchFailure(op.getLoc(),
"unknown struct width");
322 int64_t consumedBits = 0;
323 for (
size_t i = 0; i < fieldIndex; ++i) {
324 int64_t fieldWidth = hw::getBitWidth(elements[i].type);
326 "must be failed before if field width is unknown");
327 consumedBits += fieldWidth;
330 int64_t fieldWidth = hw::getBitWidth(elements[fieldIndex].type);
332 "must be failed before if field width is unknown");
335 int64_t bitOffset = totalBitWidth - consumedBits - fieldWidth;
337 bitOffset, fieldWidth);
347 ConversionPatternRewriter &rewriter)
const override {
348 auto inputTy = adaptor.getInput().getType();
349 auto outputTy = typeConverter->convertType(op.getType());
351 return rewriter.notifyMatchFailure(op,
"Failed to convert result type.");
353 auto inBits = hw::getBitWidth(inputTy);
354 auto outBits = hw::getBitWidth(outputTy);
355 if (inBits != outBits)
356 return rewriter.notifyMatchFailure(
357 op,
"Width of converted types does not match.");
359 return rewriter.notifyMatchFailure(op,
"Unknown bitwidth.");
361 auto bitcastOp = rewriter.createOrFold<
hw::BitcastOp>(op.getLoc(), outputTy,
363 rewriter.replaceOp(op, bitcastOp);
373 ConversionPatternRewriter &rewriter)
const override {
376 op, adaptor.getCond(), adaptor.getTrueValue(), adaptor.getFalseValue());
383class AggregateTypeConverter :
public TypeConverter {
385 AggregateTypeConverter() {
386 addConversion([](Type type) -> Type {
return type; });
387 addConversion([](hw::ArrayType t) -> Type {
388 auto bitWidth = t.getBitWidth();
391 return IntegerType::get(t.getContext(), *bitWidth);
393 addConversion([](hw::StructType t) -> Type {
394 auto bitWidth = t.getBitWidth();
397 return IntegerType::get(t.getContext(), *bitWidth);
399 addConversion([](hw::UnionType t) -> Type {
400 auto bitWidth = t.getBitWidth();
403 return IntegerType::get(t.getContext(), *bitWidth);
405 addTargetMaterialization([](mlir::OpBuilder &builder, mlir::Type resultType,
406 mlir::ValueRange inputs,
407 mlir::Location loc) -> mlir::Value {
408 if (inputs.size() != 1)
415 addSourceMaterialization([](mlir::OpBuilder &builder, mlir::Type resultType,
416 mlir::ValueRange inputs,
417 mlir::Location loc) -> mlir::Value {
418 if (inputs.size() != 1)
429 RewritePatternSet &
patterns, AggregateTypeConverter &typeConverter) {
431 HWArrayGetOpConversion, HWArrayCreateLikeOpConversion<hw::ArrayCreateOp>,
432 HWArrayCreateLikeOpConversion<hw::ArrayConcatOp>,
433 HWAggregateConstantOpConversion, HWArraySliceOpConversion,
434 HWArrayInjectOpConversion, HWStructCreateOpConversion,
435 HWStructExtractOpConversion, HWUnionCreateOpConversion,
436 HWUnionExtractOpConversion, BitcastOpConversion, MuxOpConversion>(
437 typeConverter,
patterns.getContext());
441struct HWAggregateToCombPass
442 :
public hw::impl::HWAggregateToCombBase<HWAggregateToCombPass> {
443 void runOnOperation()
override;
444 using HWAggregateToCombBase<HWAggregateToCombPass>::HWAggregateToCombBase;
448void HWAggregateToCombPass::runOnOperation() {
449 ConversionTarget target(getContext());
452 hw::AggregateConstantOp, hw::ArrayInjectOp,
454 hw::UnionCreateOp, hw::UnionExtractOp>();
455 target.addLegalDialect<hw::HWDialect, comb::CombDialect>();
457 RewritePatternSet
patterns(&getContext());
458 AggregateTypeConverter typeConverter;
462 [&typeConverter](
auto op) {
return typeConverter.isLegal(op); });
464 if (failed(mlir::applyPartialConversion(getOperation(), target,
466 return signalPassFailure();
assert(baseType &&"element must be base type")
MlirType uint64_t numElements
static void populateHWAggregateToCombOpConversionPatterns(RewritePatternSet &patterns, AggregateTypeConverter &typeConverter)
create(elements, Type result_type=None)
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.