CIRCT 24.0.0git
Loading...
Searching...
No Matches
HWAggregateToComb.cpp
Go to the documentation of this file.
1//===- HWAggregateToComb.cpp - HW aggregate to comb -------------*- C++ -*-===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
13#include "mlir/Pass/Pass.h"
14#include "mlir/Transforms/DialectConversion.h"
15#include "llvm/ADT/APInt.h"
16
17namespace circt {
18namespace hw {
19#define GEN_PASS_DEF_HWAGGREGATETOCOMB
20#include "circt/Dialect/HW/Passes.h.inc"
21} // namespace hw
22} // namespace circt
23
24using namespace mlir;
25using namespace circt;
26
27namespace {
28
29// Lower hw.array_create and hw.array_concat to comb.concat.
30template <typename OpTy>
31struct HWArrayCreateLikeOpConversion : OpConversionPattern<OpTy> {
33 using OpAdaptor = typename OpConversionPattern<OpTy>::OpAdaptor;
34 LogicalResult
35 matchAndRewrite(OpTy op, OpAdaptor adaptor,
36 ConversionPatternRewriter &rewriter) const override {
37 rewriter.replaceOpWithNewOp<comb::ConcatOp>(op, adaptor.getInputs());
38 return success();
39 }
40};
41
42struct HWUnionCreateOpConversion
43 : public OpConversionPattern<hw::UnionCreateOp> {
44 using OpConversionPattern<hw::UnionCreateOp>::OpConversionPattern;
45 // hw.union_create -> hw.bitcast [ + comb.concat ]
46 LogicalResult
47 matchAndRewrite(hw::UnionCreateOp op, OpAdaptor adaptor,
48 ConversionPatternRewriter &rewriter) const override {
49 hw::UnionType unionTy = op.getType();
50 auto outputTy =
51 dyn_cast_or_null<IntegerType>(typeConverter->convertType(unionTy));
52 if (!outputTy)
53 return rewriter.notifyMatchFailure(op.getLoc(),
54 "Failed to convert union to integer");
55
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");
60
61 // Bitcast the input value to its integer representation.
62 auto inputIntTy = rewriter.getIntegerType(inputBitWidth);
63 Value inputAsInt = rewriter.createOrFold<hw::BitcastOp>(
64 op.getLoc(), inputIntTy, adaptor.getInput());
65
66 // The field shares the LSB of the union and is moved towards the MSB by
67 // its offset. The bits the field does not cover are undefined and filled
68 // with zeros.
69 int64_t bitOffset = unionTy.getElements()[op.getFieldIndex()].offset;
70 int64_t prePadding = outputTy.getWidth() - inputBitWidth - bitOffset;
71
72 auto createZeroCst = [&](Location loc, int64_t bitWidth) -> Value {
73 return hw::ConstantOp::create(rewriter, loc,
74 rewriter.getIntegerType(bitWidth), 0);
75 };
76
77 SmallVector<Value> concatOperands;
78 if (prePadding > 0)
79 concatOperands.push_back(createZeroCst(op.getLoc(), prePadding));
80 concatOperands.push_back(inputAsInt);
81 if (bitOffset > 0)
82 concatOperands.push_back(createZeroCst(op.getLoc(), bitOffset));
83
84 Value result =
85 rewriter.createOrFold<comb::ConcatOp>(op.getLoc(), concatOperands);
86 rewriter.replaceOp(op, result);
87 return success();
88 }
89};
90
91struct HWUnionExtractOpConversion
92 : public OpConversionPattern<hw::UnionExtractOp> {
93 using OpConversionPattern<hw::UnionExtractOp>::OpConversionPattern;
94 // hw.union_extract -> [ comb.extract + ] hw.bitcast
95 LogicalResult
96 matchAndRewrite(hw::UnionExtractOp op, OpAdaptor adaptor,
97 ConversionPatternRewriter &rewriter) const override {
98 hw::UnionType unionTy = op.getInput().getType();
99
100 auto inputTy = dyn_cast_or_null<IntegerType>(adaptor.getInput().getType());
101 if (!inputTy)
102 return rewriter.notifyMatchFailure(op.getLoc(),
103 "Failed to convert union to integer");
104 auto outputTy = typeConverter->convertType(op.getType());
105 if (!outputTy)
106 return rewriter.notifyMatchFailure(
107 op.getLoc(), "Failed to convert union extract result type");
108
109 auto resultFieldBits = hw::getBitWidth(outputTy);
110 assert(resultFieldBits >= 0);
111 auto integerValue = adaptor.getInput();
112
113 // If the output is narrower than the union, extract the active bits.
114 if (resultFieldBits < integerValue.getType().getIntOrFloatBitWidth()) {
115 auto bitOffset = unionTy.getElements()[op.getFieldIndex()].offset;
116 integerValue = comb::ExtractOp::create(
117 rewriter, op->getLoc(), rewriter.getIntegerType(resultFieldBits),
118 integerValue, bitOffset);
119 }
120
121 // Bitcast the extracted bits to the result. Fold inplace if outputTy ==
122 // inputTy.
123 auto bitcastOp = rewriter.createOrFold<hw::BitcastOp>(op.getLoc(), outputTy,
124 integerValue);
125
126 rewriter.replaceOp(op, bitcastOp);
127 return success();
128 }
129};
130
131struct HWAggregateConstantOpConversion
132 : OpConversionPattern<hw::AggregateConstantOp> {
133 using OpConversionPattern<hw::AggregateConstantOp>::OpConversionPattern;
134
135 LogicalResult
136 matchAndRewrite(hw::AggregateConstantOp op, OpAdaptor adaptor,
137 ConversionPatternRewriter &rewriter) const override {
138 // Lower to concat.
139 APInt intVal;
140 if (failed(hw::aggregateAttrToAPInt(op.getType(), adaptor.getFieldsAttr(),
141 intVal)))
142 return failure();
143 rewriter.replaceOpWithNewOp<hw::ConstantOp>(op, intVal);
144 return success();
145 }
146};
147
148struct HWArrayGetOpConversion : OpConversionPattern<hw::ArrayGetOp> {
150
151 LogicalResult
152 matchAndRewrite(hw::ArrayGetOp op, OpAdaptor adaptor,
153 ConversionPatternRewriter &rewriter) const override {
154 SmallVector<Value> results;
155 auto arrayType = cast<hw::ArrayType>(op.getInput().getType());
156 auto elemType = arrayType.getElementType();
157 auto numElements = arrayType.getNumElements();
158 auto elemWidth = hw::getBitWidth(elemType);
159 if (elemWidth < 0)
160 return rewriter.notifyMatchFailure(op.getLoc(), "unknown element width");
161
162 auto lowered = adaptor.getInput();
163 auto index = adaptor.getIndex();
164 APInt constantIndex;
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)) {
169 rewriter.replaceOpWithNewOp<comb::ExtractOp>(
170 op, lowered, constantIndex.getZExtValue() * elemWidth, elemWidth);
171 return success();
172 }
173 }
174
175 for (size_t i = 0; i < numElements; ++i)
176 results.push_back(rewriter.createOrFold<comb::ExtractOp>(
177 op.getLoc(), lowered, i * elemWidth, elemWidth));
178
179 SmallVector<Value> bits;
180 comb::extractBits(rewriter, index, bits);
181 auto result = comb::constructMuxTree(rewriter, op.getLoc(), bits, results,
182 results.back());
183
184 rewriter.replaceOp(op, result);
185 return success();
186 }
187};
188
189struct HWArraySliceOpConversion : OpConversionPattern<hw::ArraySliceOp> {
191
192 LogicalResult
193 matchAndRewrite(hw::ArraySliceOp op, OpAdaptor adaptor,
194 ConversionPatternRewriter &rewriter) const override {
195 SmallVector<Value> results;
196 auto arrayType = cast<hw::ArrayType>(op.getInput().getType());
197 auto elemType = arrayType.getElementType();
198 auto numElements = arrayType.getNumElements();
199 auto elemWidth = hw::getBitWidth(elemType);
200 if (elemWidth < 0)
201 return rewriter.notifyMatchFailure(op.getLoc(), "unknown element width");
202 auto resultArrayType = cast<hw::ArrayType>(op.getResult().getType());
203 auto resultNumElements = resultArrayType.getNumElements();
204
205 auto lowered = adaptor.getInput();
206 auto index = adaptor.getLowIndex();
207 APInt constantIndex;
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)) {
212 rewriter.replaceOpWithNewOp<comb::ExtractOp>(
213 op, lowered, constantIndex.getZExtValue() * elemWidth,
214 resultNumElements * elemWidth);
215 return success();
216 }
217 }
218
219 for (size_t i = 0; i <= numElements - resultNumElements; ++i)
220 results.push_back(rewriter.createOrFold<comb::ExtractOp>(
221 op.getLoc(), lowered, i * elemWidth, resultNumElements * elemWidth));
222
223 SmallVector<Value> bits;
224 comb::extractBits(rewriter, index, bits);
225 auto result = comb::constructMuxTree(rewriter, op.getLoc(), bits, results,
226 results.back());
227
228 rewriter.replaceOp(op, result);
229 return success();
230 }
231};
232
233struct HWArrayInjectOpConversion : OpConversionPattern<hw::ArrayInjectOp> {
234 using OpConversionPattern<hw::ArrayInjectOp>::OpConversionPattern;
235
236 LogicalResult
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();
241 auto numElements = arrayType.getNumElements();
242 auto elemWidth = hw::getBitWidth(elemType);
243 if (elemWidth < 0)
244 return rewriter.notifyMatchFailure(op.getLoc(), "unknown element width");
245
246 Location loc = op.getLoc();
247
248 // Extract all elements from the input array
249 SmallVector<Value> originalElements;
250 auto inputArray = adaptor.getInput();
251 for (size_t i = 0; i < numElements; ++i) {
252 originalElements.push_back(rewriter.createOrFold<comb::ExtractOp>(
253 loc, inputArray, i * elemWidth, elemWidth));
254 }
255
256 // Create 2D array: each row represents what the array would look like
257 // if injection happened at that specific index
258 SmallVector<Value> arrayRows;
259 arrayRows.reserve(numElements);
260 for (int injectIdx = numElements - 1; injectIdx >= 0; --injectIdx) {
261 SmallVector<Value> rowElements;
262 rowElements.reserve(numElements);
263
264 // Build the row: array[n-1], array[n-2], ..., but replace element at
265 // injectIdx with newVal
266 for (int originalIdx = numElements - 1; originalIdx >= 0; --originalIdx) {
267 if (originalIdx == injectIdx) {
268 rowElements.push_back(adaptor.getElement());
269 } else {
270 rowElements.push_back(originalElements[originalIdx]);
271 }
272 }
273
274 // Concatenate elements to form this row
275 Value row = hw::ArrayCreateOp::create(rewriter, loc, rowElements);
276 arrayRows.push_back(row);
277 }
278
279 // Create the 2D array by concatenating all rows
280 // arrayRows[0] corresponds to injection at index 0
281 // arrayRows[1] corresponds to injection at index 1, etc.
282 Value array2D = hw::ArrayCreateOp::create(rewriter, loc, arrayRows);
283
284 // Create array_get operation to select the row
285 auto arrayGetOp =
286 hw::ArrayGetOp::create(rewriter, loc, array2D, adaptor.getIndex());
287
288 rewriter.replaceOp(op, arrayGetOp);
289 return success();
290 }
291};
292
293struct HWStructCreateOpConversion : OpConversionPattern<hw::StructCreateOp> {
295
296 LogicalResult
297 matchAndRewrite(hw::StructCreateOp op, OpAdaptor adaptor,
298 ConversionPatternRewriter &rewriter) const override {
299 // Lower struct_create to comb.concat. The first field occupies the MSBs, so
300 // we concatenate fields in order (comb.concat places first operand at MSB).
301 rewriter.replaceOpWithNewOp<comb::ConcatOp>(op, adaptor.getInput());
302 return success();
303 }
304};
305
306struct HWStructExtractOpConversion : OpConversionPattern<hw::StructExtractOp> {
308
309 LogicalResult
310 matchAndRewrite(hw::StructExtractOp op, OpAdaptor adaptor,
311 ConversionPatternRewriter &rewriter) const override {
312 auto structType = cast<hw::StructType>(op.getInput().getType());
313 auto fieldIndex = op.getFieldIndex();
314 auto elements = structType.getElements();
315
316 int64_t totalBitWidth = hw::getBitWidth(structType);
317 if (totalBitWidth < 0)
318 return rewriter.notifyMatchFailure(op.getLoc(), "unknown struct width");
319
320 // Compute the bit offset from the MSB by summing the widths of all
321 // preceding fields. The first field occupies the MSBs.
322 int64_t consumedBits = 0;
323 for (size_t i = 0; i < fieldIndex; ++i) {
324 int64_t fieldWidth = hw::getBitWidth(elements[i].type);
325 assert(fieldWidth >= 0 &&
326 "must be failed before if field width is unknown");
327 consumedBits += fieldWidth;
328 }
329
330 int64_t fieldWidth = hw::getBitWidth(elements[fieldIndex].type);
331 assert(fieldWidth >= 0 &&
332 "must be failed before if field width is unknown");
333
334 // Extract the field using comb.extract. Offset is from LSB.
335 int64_t bitOffset = totalBitWidth - consumedBits - fieldWidth;
336 rewriter.replaceOpWithNewOp<comb::ExtractOp>(op, adaptor.getInput(),
337 bitOffset, fieldWidth);
338 return success();
339 }
340};
341
342struct BitcastOpConversion : OpConversionPattern<hw::BitcastOp> {
344 // Recreate bitcast with legalized types.
345 LogicalResult
346 matchAndRewrite(hw::BitcastOp op, OpAdaptor adaptor,
347 ConversionPatternRewriter &rewriter) const override {
348 auto inputTy = adaptor.getInput().getType();
349 auto outputTy = typeConverter->convertType(op.getType());
350 if (!outputTy)
351 return rewriter.notifyMatchFailure(op, "Failed to convert result type.");
352
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.");
358 if (inBits < 0)
359 return rewriter.notifyMatchFailure(op, "Unknown bitwidth.");
360
361 auto bitcastOp = rewriter.createOrFold<hw::BitcastOp>(op.getLoc(), outputTy,
362 adaptor.getInput());
363 rewriter.replaceOp(op, bitcastOp);
364 return success();
365 }
366};
367
368struct MuxOpConversion : OpConversionPattern<comb::MuxOp> {
370
371 LogicalResult
372 matchAndRewrite(comb::MuxOp op, OpAdaptor adaptor,
373 ConversionPatternRewriter &rewriter) const override {
374 // Re-create Mux with legalized types.
375 rewriter.replaceOpWithNewOp<comb::MuxOp>(
376 op, adaptor.getCond(), adaptor.getTrueValue(), adaptor.getFalseValue());
377 return success();
378 }
379};
380
381/// A type converter is needed to perform the in-flight materialization of
382/// aggregate types to integer types.
383class AggregateTypeConverter : public TypeConverter {
384public:
385 AggregateTypeConverter() {
386 addConversion([](Type type) -> Type { return type; });
387 addConversion([](hw::ArrayType t) -> Type {
388 auto bitWidth = t.getBitWidth();
389 if (!bitWidth)
390 return {};
391 return IntegerType::get(t.getContext(), *bitWidth);
392 });
393 addConversion([](hw::StructType t) -> Type {
394 auto bitWidth = t.getBitWidth();
395 if (!bitWidth)
396 return {};
397 return IntegerType::get(t.getContext(), *bitWidth);
398 });
399 addConversion([](hw::UnionType t) -> Type {
400 auto bitWidth = t.getBitWidth();
401 if (!bitWidth)
402 return {};
403 return IntegerType::get(t.getContext(), *bitWidth);
404 });
405 addTargetMaterialization([](mlir::OpBuilder &builder, mlir::Type resultType,
406 mlir::ValueRange inputs,
407 mlir::Location loc) -> mlir::Value {
408 if (inputs.size() != 1)
409 return Value();
410
411 return hw::BitcastOp::create(builder, loc, resultType, inputs[0])
412 ->getResult(0);
413 });
414
415 addSourceMaterialization([](mlir::OpBuilder &builder, mlir::Type resultType,
416 mlir::ValueRange inputs,
417 mlir::Location loc) -> mlir::Value {
418 if (inputs.size() != 1)
419 return Value();
420
421 return hw::BitcastOp::create(builder, loc, resultType, inputs[0])
422 ->getResult(0);
423 });
424 }
425};
426} // namespace
427
429 RewritePatternSet &patterns, AggregateTypeConverter &typeConverter) {
430 patterns.add<
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());
438}
439
440namespace {
441struct HWAggregateToCombPass
442 : public hw::impl::HWAggregateToCombBase<HWAggregateToCombPass> {
443 void runOnOperation() override;
444 using HWAggregateToCombBase<HWAggregateToCombPass>::HWAggregateToCombBase;
445};
446} // namespace
447
448void HWAggregateToCombPass::runOnOperation() {
449 ConversionTarget target(getContext());
450
452 hw::AggregateConstantOp, hw::ArrayInjectOp,
454 hw::UnionCreateOp, hw::UnionExtractOp>();
455 target.addLegalDialect<hw::HWDialect, comb::CombDialect>();
456
457 RewritePatternSet patterns(&getContext());
458 AggregateTypeConverter typeConverter;
460
461 target.addDynamicallyLegalOp<comb::MuxOp, hw::BitcastOp>(
462 [&typeConverter](auto op) { return typeConverter.isLegal(op); });
463
464 if (failed(mlir::applyPartialConversion(getOperation(), target,
465 std::move(patterns))))
466 return signalPassFailure();
467}
assert(baseType &&"element must be base type")
MlirType uint64_t numElements
Definition CHIRRTL.cpp:30
static void populateHWAggregateToCombOpConversionPatterns(RewritePatternSet &patterns, AggregateTypeConverter &typeConverter)
create(low_bit, result_type, input=None)
Definition comb.py:187
create(elements, Type result_type=None)
Definition hw.py:483
create(array_value, idx)
Definition hw.py:450
create(data_type, value)
Definition hw.py:441
create(data_type, value)
Definition hw.py:433
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.
Definition hw.py:1