12#include "mlir/Pass/Pass.h"
13#include "mlir/Transforms/DialectConversion.h"
14#include "llvm/ADT/APInt.h"
18#define GEN_PASS_DEF_HWAGGREGATETOCOMB
19#include "circt/Dialect/HW/Passes.h.inc"
29template <
typename OpTy>
34 matchAndRewrite(OpTy op, OpAdaptor adaptor,
35 ConversionPatternRewriter &rewriter)
const override {
36 rewriter.replaceOpWithNewOp<
comb::ConcatOp>(op, adaptor.getInputs());
41struct HWAggregateConstantOpConversion
46 matchAndRewrite(hw::AggregateConstantOp op, OpAdaptor adaptor,
47 ConversionPatternRewriter &rewriter)
const override {
50 if (failed(hw::aggregateAttrToAPInt(op.getType(), adaptor.getFieldsAttr(),
63 ConversionPatternRewriter &rewriter)
const override {
64 SmallVector<Value> results;
65 auto arrayType = cast<hw::ArrayType>(op.getInput().getType());
66 auto elemType = arrayType.getElementType();
68 auto elemWidth = hw::getBitWidth(elemType);
70 return rewriter.notifyMatchFailure(op.getLoc(),
"unknown element width");
72 auto lowered = adaptor.getInput();
73 auto index = adaptor.getIndex();
75 if (matchPattern(index, m_ConstantInt(&constantIndex))) {
76 int64_t maxIndex = std::numeric_limits<int32_t>::max() / elemWidth;
77 if (constantIndex.isSingleWord() &&
78 constantIndex.getZExtValue() <=
static_cast<uint64_t
>(maxIndex)) {
80 op, lowered, constantIndex.getZExtValue() * elemWidth, elemWidth);
87 op.getLoc(), lowered, i * elemWidth, elemWidth));
89 SmallVector<Value> bits;
90 comb::extractBits(rewriter, index, bits);
91 auto result = comb::constructMuxTree(rewriter, op.getLoc(), bits, results,
94 rewriter.replaceOp(op, result);
104 ConversionPatternRewriter &rewriter)
const override {
105 SmallVector<Value> results;
106 auto arrayType = cast<hw::ArrayType>(op.getInput().getType());
107 auto elemType = arrayType.getElementType();
109 auto elemWidth = hw::getBitWidth(elemType);
111 return rewriter.notifyMatchFailure(op.getLoc(),
"unknown element width");
112 auto resultArrayType = cast<hw::ArrayType>(op.getResult().getType());
113 auto resultNumElements = resultArrayType.getNumElements();
115 auto lowered = adaptor.getInput();
116 auto index = adaptor.getLowIndex();
118 if (matchPattern(index, m_ConstantInt(&constantIndex))) {
119 int64_t maxIndex = std::numeric_limits<int32_t>::max() / elemWidth;
120 if (constantIndex.isSingleWord() &&
121 constantIndex.getZExtValue() <=
static_cast<uint64_t
>(maxIndex)) {
123 op, lowered, constantIndex.getZExtValue() * elemWidth,
124 resultNumElements * elemWidth);
129 for (
size_t i = 0; i <=
numElements - resultNumElements; ++i)
131 op.getLoc(), lowered, i * elemWidth, resultNumElements * elemWidth));
133 SmallVector<Value> bits;
134 comb::extractBits(rewriter, index, bits);
135 auto result = comb::constructMuxTree(rewriter, op.getLoc(), bits, results,
138 rewriter.replaceOp(op, result);
147 matchAndRewrite(hw::ArrayInjectOp op, OpAdaptor adaptor,
148 ConversionPatternRewriter &rewriter)
const override {
149 auto arrayType = cast<hw::ArrayType>(op.getInput().getType());
150 auto elemType = arrayType.getElementType();
152 auto elemWidth = hw::getBitWidth(elemType);
154 return rewriter.notifyMatchFailure(op.getLoc(),
"unknown element width");
156 Location loc = op.getLoc();
159 SmallVector<Value> originalElements;
160 auto inputArray = adaptor.getInput();
163 loc, inputArray, i * elemWidth, elemWidth));
168 SmallVector<Value> arrayRows;
170 for (
int injectIdx =
numElements - 1; injectIdx >= 0; --injectIdx) {
171 SmallVector<Value> rowElements;
176 for (
int originalIdx =
numElements - 1; originalIdx >= 0; --originalIdx) {
177 if (originalIdx == injectIdx) {
178 rowElements.push_back(adaptor.getElement());
180 rowElements.push_back(originalElements[originalIdx]);
186 arrayRows.push_back(row);
198 rewriter.replaceOp(op, arrayGetOp);
208 ConversionPatternRewriter &rewriter)
const override {
211 rewriter.replaceOpWithNewOp<
comb::ConcatOp>(op, adaptor.getInput());
221 ConversionPatternRewriter &rewriter)
const override {
222 auto structType = cast<hw::StructType>(op.getInput().getType());
223 auto fieldIndex = op.getFieldIndex();
224 auto elements = structType.getElements();
226 int64_t totalBitWidth = hw::getBitWidth(structType);
227 if (totalBitWidth < 0)
228 return rewriter.notifyMatchFailure(op.getLoc(),
"unknown struct width");
232 int64_t consumedBits = 0;
233 for (
size_t i = 0; i < fieldIndex; ++i) {
234 int64_t fieldWidth = hw::getBitWidth(elements[i].type);
236 "must be failed before if field width is unknown");
237 consumedBits += fieldWidth;
240 int64_t fieldWidth = hw::getBitWidth(elements[fieldIndex].type);
242 "must be failed before if field width is unknown");
245 int64_t bitOffset = totalBitWidth - consumedBits - fieldWidth;
247 bitOffset, fieldWidth);
257 ConversionPatternRewriter &rewriter)
const override {
260 op, adaptor.getCond(), adaptor.getTrueValue(), adaptor.getFalseValue());
267class AggregateTypeConverter :
public TypeConverter {
269 AggregateTypeConverter() {
270 addConversion([](Type type) -> Type {
return type; });
271 addConversion([](hw::ArrayType t) -> Type {
272 return IntegerType::get(t.getContext(), hw::getBitWidth(t));
274 addConversion([](hw::StructType t) -> Type {
275 return IntegerType::get(t.getContext(), hw::getBitWidth(t));
277 addTargetMaterialization([](mlir::OpBuilder &builder, mlir::Type resultType,
278 mlir::ValueRange inputs,
279 mlir::Location loc) -> mlir::Value {
280 if (inputs.size() != 1)
287 addSourceMaterialization([](mlir::OpBuilder &builder, mlir::Type resultType,
288 mlir::ValueRange inputs,
289 mlir::Location loc) -> mlir::Value {
290 if (inputs.size() != 1)
301 RewritePatternSet &
patterns, AggregateTypeConverter &typeConverter) {
302 patterns.add<HWArrayGetOpConversion,
303 HWArrayCreateLikeOpConversion<hw::ArrayCreateOp>,
304 HWArrayCreateLikeOpConversion<hw::ArrayConcatOp>,
305 HWAggregateConstantOpConversion, HWArraySliceOpConversion,
306 HWArrayInjectOpConversion, HWStructCreateOpConversion,
307 HWStructExtractOpConversion, MuxOpConversion>(
308 typeConverter,
patterns.getContext());
312struct HWAggregateToCombPass
313 :
public hw::impl::HWAggregateToCombBase<HWAggregateToCombPass> {
314 void runOnOperation()
override;
315 using HWAggregateToCombBase<HWAggregateToCombPass>::HWAggregateToCombBase;
319void HWAggregateToCombPass::runOnOperation() {
320 ConversionTarget target(getContext());
323 hw::AggregateConstantOp, hw::ArrayInjectOp,
327 [](
comb::MuxOp op) {
return hw::type_isa<IntegerType>(op.getType()); });
328 target.addLegalDialect<hw::HWDialect, comb::CombDialect>();
330 RewritePatternSet
patterns(&getContext());
331 AggregateTypeConverter typeConverter;
334 if (failed(mlir::applyPartialConversion(getOperation(), target,
336 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.