16#include "mlir/IR/Builders.h"
17#include "mlir/IR/ImplicitLocOpBuilder.h"
18#include "mlir/IR/Matchers.h"
19#include "mlir/IR/PatternMatch.h"
20#include "llvm/Support/FormatVariadic.h"
26using namespace matchers;
32bool comb::shouldUseBoothEncoding(Value lhs, Value rhs,
unsigned threshold) {
38 auto lhsWidth = lhs.getType().getIntOrFloatBitWidth();
39 auto rhsWidth = rhs.getType().getIntOrFloatBitWidth();
42 Value lhsZext, rhsZext;
43 if (matchPattern(lhs, comb::m_ZextBy(m_Any(&lhsZext))))
44 lhsWidth -= lhsZext.getType().getIntOrFloatBitWidth();
45 if (matchPattern(rhs, comb::m_ZextBy(m_Any(&rhsZext))))
46 rhsWidth -= rhsZext.getType().getIntOrFloatBitWidth();
49 Value lhsSextBits, rhsSextBits;
50 if (matchPattern(lhs, comb::m_SextBy(m_Any(&lhsSextBits))))
51 lhsWidth -= lhsSextBits.getType().getIntOrFloatBitWidth();
52 if (matchPattern(rhs, comb::m_SextBy(m_Any(&rhsSextBits))))
53 rhsWidth -= rhsSextBits.getType().getIntOrFloatBitWidth();
58 return lhsWidth > threshold && rhsWidth > threshold;
61Value comb::createZExt(OpBuilder &builder, Location loc, Value value,
62 unsigned targetWidth) {
63 assert(value.getType().isSignlessInteger());
64 auto inputWidth = value.getType().getIntOrFloatBitWidth();
65 assert(inputWidth <= targetWidth);
68 if (inputWidth == targetWidth)
73 builder, loc, builder.getIntegerType(targetWidth - inputWidth), 0);
74 return builder.createOrFold<
ConcatOp>(loc, zeros, value);
79Value comb::createOrFoldSExt(OpBuilder &builder, Location loc, Value value,
81 IntegerType valueType = dyn_cast<IntegerType>(value.getType());
82 assert(valueType && isa<IntegerType>(destTy) &&
83 valueType.getWidth() <= destTy.getIntOrFloatBitWidth() &&
84 valueType.getWidth() != 0 &&
"invalid sext operands");
86 if (valueType == destTy)
91 builder.createOrFold<
ExtractOp>(loc, value, valueType.getWidth() - 1, 1);
92 auto signBits = builder.createOrFold<ReplicateOp>(
93 loc, signBit, destTy.getIntOrFloatBitWidth() - valueType.getWidth());
94 return builder.createOrFold<
ConcatOp>(loc, signBits, value);
97Value comb::createOrFoldSExt(ImplicitLocOpBuilder &builder, Value value,
102Value comb::createOrFoldNot(OpBuilder &builder, Location loc, Value value,
105 return builder.createOrFold<
XorOp>(loc, value, allOnes, twoState);
108Value comb::createOrFoldNot(ImplicitLocOpBuilder &builder, Value value,
114void comb::extractBits(OpBuilder &builder, Value val,
115 SmallVectorImpl<Value> &bits) {
116 assert(val.getType().isInteger() &&
"expected integer");
117 auto width = val.getType().getIntOrFloatBitWidth();
122 if (concat.getNumOperands() == width &&
123 llvm::all_of(concat.getOperandTypes(), [](Type type) {
124 return type.getIntOrFloatBitWidth() == 1;
127 bits.append(std::make_reverse_iterator(concat.getOperands().end()),
128 std::make_reverse_iterator(concat.getOperands().begin()));
134 for (int64_t i = 0; i < width; ++i)
141Value comb::constructMuxTree(OpBuilder &builder, Location loc,
142 ArrayRef<Value> selectors,
143 ArrayRef<Value> leafNodes,
144 Value outOfBoundsValue) {
146 std::function<Value(
size_t,
size_t)> constructTreeHelper =
147 [&](
size_t id,
size_t level) -> Value {
152 return id < leafNodes.size() ? leafNodes[id] : outOfBoundsValue;
155 auto selector = selectors[level - 1];
158 auto trueVal = constructTreeHelper(2 *
id + 1, level - 1);
159 auto falseVal = constructTreeHelper(2 *
id, level - 1);
162 return builder.createOrFold<
comb::MuxOp>(loc, selector, trueVal, falseVal);
165 return constructTreeHelper(0, llvm::Log2_64_Ceil(leafNodes.size()));
168Value comb::createDynamicExtract(OpBuilder &builder, Location loc, Value value,
169 Value offset,
unsigned width) {
170 assert(value.getType().isSignlessInteger());
171 auto valueWidth = value.getType().getIntOrFloatBitWidth();
172 assert(width <= valueWidth);
176 if (matchPattern(offset, mlir::m_ConstantInt(&constOffset)))
177 if (constOffset.getActiveBits() < 32)
179 loc, value, constOffset.getZExtValue(), width);
183 offset =
createZExt(builder, loc, offset, valueWidth);
184 value = builder.createOrFold<
comb::ShrUOp>(loc, value, offset);
188Value comb::createDynamicInject(OpBuilder &builder, Location loc, Value value,
189 Value offset, Value replacement,
191 assert(value.getType().isSignlessInteger());
192 assert(replacement.getType().isSignlessInteger());
193 auto largeWidth = value.getType().getIntOrFloatBitWidth();
194 auto smallWidth = replacement.getType().getIntOrFloatBitWidth();
195 assert(smallWidth <= largeWidth);
203 if (matchPattern(offset, mlir::m_ConstantInt(&constOffset)))
204 if (constOffset.getActiveBits() < 32)
205 return createInject(builder, loc, value, constOffset.getZExtValue(),
209 offset =
createZExt(builder, loc, offset, largeWidth);
211 builder, loc, APInt::getLowBitsSet(largeWidth, smallWidth));
212 mask = builder.createOrFold<
comb::ShlOp>(loc, mask, offset);
214 value = builder.createOrFold<
comb::AndOp>(loc, value, mask, twoState);
218 replacement =
createZExt(builder, loc, replacement, largeWidth);
219 replacement = builder.createOrFold<
comb::ShlOp>(loc, replacement, offset);
220 return builder.createOrFold<
comb::OrOp>(loc, value, replacement, twoState);
223Value comb::createInject(OpBuilder &builder, Location loc, Value value,
224 unsigned offset, Value replacement) {
225 assert(value.getType().isSignlessInteger());
226 assert(replacement.getType().isSignlessInteger());
227 auto largeWidth = value.getType().getIntOrFloatBitWidth();
228 auto smallWidth = replacement.getType().getIntOrFloatBitWidth();
229 assert(smallWidth <= largeWidth);
232 if (offset >= largeWidth)
241 SmallVector<Value, 3> fragments;
242 auto end = offset + smallWidth;
243 if (end < largeWidth)
246 if (end <= largeWidth)
247 fragments.push_back(replacement);
250 largeWidth - offset));
258 mlir::PatternRewriter &rewriter) {
259 auto lhs = subOp.getLhs();
260 auto rhs = subOp.getRhs();
265 comb::createOrFoldNot(rewriter, subOp.getLoc(), rhs, subOp.getTwoState());
268 replaceOpWithNewOpAndCopyNamehint<comb::AddOp>(
269 rewriter, subOp, ValueRange{lhs, notRhs, one}, subOp.getTwoState());
274 Operation *op, Value lhs,
275 Value rhs,
bool isDiv) {
281 APInt rhsValue = rhsConstantOp.getValue();
282 if (!rhsValue.isPowerOf2())
285 Location loc = op->getLoc();
287 unsigned width = lhs.getType().getIntOrFloatBitWidth();
288 unsigned bitPosition = rhsValue.ceilLogBase2();
296 loc, lhs, bitPosition, width - bitPosition);
305 comb::ConcatOp::create(rewriter, loc,
306 ArrayRef<Value>{zeros, upperBits}));
319 APInt::getZero(width - bitPosition));
323 comb::ConcatOp::create(rewriter, loc, ArrayRef<Value>{zeros, lowerBits}));
327LogicalResult comb::convertDivUByPowerOfTwo(
DivUOp divOp,
328 mlir::PatternRewriter &rewriter) {
330 divOp.getRhs(),
true);
333LogicalResult comb::convertModUByPowerOfTwo(
ModUOp modOp,
334 mlir::PatternRewriter &rewriter) {
336 modOp.getRhs(),
false);
343ICmpPredicate ICmpOp::getFlippedPredicate(ICmpPredicate predicate) {
345 case ICmpPredicate::eq:
346 return ICmpPredicate::eq;
347 case ICmpPredicate::ne:
348 return ICmpPredicate::ne;
349 case ICmpPredicate::slt:
350 return ICmpPredicate::sgt;
351 case ICmpPredicate::sle:
352 return ICmpPredicate::sge;
353 case ICmpPredicate::sgt:
354 return ICmpPredicate::slt;
355 case ICmpPredicate::sge:
356 return ICmpPredicate::sle;
357 case ICmpPredicate::ult:
358 return ICmpPredicate::ugt;
359 case ICmpPredicate::ule:
360 return ICmpPredicate::uge;
361 case ICmpPredicate::ugt:
362 return ICmpPredicate::ult;
363 case ICmpPredicate::uge:
364 return ICmpPredicate::ule;
365 case ICmpPredicate::ceq:
366 return ICmpPredicate::ceq;
367 case ICmpPredicate::cne:
368 return ICmpPredicate::cne;
369 case ICmpPredicate::weq:
370 return ICmpPredicate::weq;
371 case ICmpPredicate::wne:
372 return ICmpPredicate::wne;
374 llvm_unreachable(
"unknown comparison predicate");
377bool ICmpOp::isPredicateSigned(ICmpPredicate predicate) {
379 case ICmpPredicate::ult:
380 case ICmpPredicate::ugt:
381 case ICmpPredicate::ule:
382 case ICmpPredicate::uge:
383 case ICmpPredicate::ne:
384 case ICmpPredicate::eq:
385 case ICmpPredicate::cne:
386 case ICmpPredicate::ceq:
387 case ICmpPredicate::wne:
388 case ICmpPredicate::weq:
390 case ICmpPredicate::slt:
391 case ICmpPredicate::sgt:
392 case ICmpPredicate::sle:
393 case ICmpPredicate::sge:
396 llvm_unreachable(
"unknown comparison predicate");
401ICmpPredicate ICmpOp::getNegatedPredicate(ICmpPredicate predicate) {
403 case ICmpPredicate::eq:
404 return ICmpPredicate::ne;
405 case ICmpPredicate::ne:
406 return ICmpPredicate::eq;
407 case ICmpPredicate::slt:
408 return ICmpPredicate::sge;
409 case ICmpPredicate::sle:
410 return ICmpPredicate::sgt;
411 case ICmpPredicate::sgt:
412 return ICmpPredicate::sle;
413 case ICmpPredicate::sge:
414 return ICmpPredicate::slt;
415 case ICmpPredicate::ult:
416 return ICmpPredicate::uge;
417 case ICmpPredicate::ule:
418 return ICmpPredicate::ugt;
419 case ICmpPredicate::ugt:
420 return ICmpPredicate::ule;
421 case ICmpPredicate::uge:
422 return ICmpPredicate::ult;
423 case ICmpPredicate::ceq:
424 return ICmpPredicate::cne;
425 case ICmpPredicate::cne:
426 return ICmpPredicate::ceq;
427 case ICmpPredicate::weq:
428 return ICmpPredicate::wne;
429 case ICmpPredicate::wne:
430 return ICmpPredicate::weq;
432 llvm_unreachable(
"unknown comparison predicate");
437bool ICmpOp::isEqualAllOnes() {
438 if (getPredicate() != ICmpPredicate::eq)
442 dyn_cast_or_null<hw::ConstantOp>(getOperand(1).getDefiningOp()))
443 return op1.getValue().isAllOnes();
449bool ICmpOp::isNotEqualZero() {
450 if (getPredicate() != ICmpPredicate::ne)
454 dyn_cast_or_null<hw::ConstantOp>(getOperand(1).getDefiningOp()))
455 return op1.getValue().isZero();
463LogicalResult ReplicateOp::verify() {
466 auto srcWidth = cast<IntegerType>(getOperand().getType()).getWidth();
467 auto dstWidth = cast<IntegerType>(getType()).getWidth();
469 if (srcWidth > dstWidth)
470 return emitOpError(
"replicate cannot shrink bitwidth of operand"),
473 if ((srcWidth == 0 && dstWidth != 0) ||
474 (srcWidth != 0 && dstWidth % srcWidth))
475 return emitOpError(
"replicate must produce integer multiple of operand"),
486 if (op->getOperands().empty())
487 return op->emitOpError(
"requires 1 or more args");
491LogicalResult AddOp::verify() {
return verifyUTBinOp(*
this); }
493LogicalResult MulOp::verify() {
return verifyUTBinOp(*
this); }
495LogicalResult AndOp::verify() {
return verifyUTBinOp(*
this); }
499LogicalResult XorOp::verify() {
return verifyUTBinOp(*
this); }
503bool XorOp::isBinaryNot() {
504 if (getNumOperands() != 2)
506 if (
auto cst = getOperand(1).getDefiningOp<hw::ConstantOp>())
507 if (cst.getValue().isAllOnes())
517 unsigned resultWidth = 0;
518 for (
auto input : inputs) {
519 resultWidth += hw::type_cast<IntegerType>(input.getType()).getWidth();
524void ConcatOp::build(OpBuilder &builder, OperationState &result, Value hd,
526 result.addOperands(ValueRange{hd});
527 result.addOperands(tl);
528 unsigned hdWidth = cast<IntegerType>(hd.getType()).getWidth();
529 result.addTypes(builder.getIntegerType(
getTotalWidth(tl) + hdWidth));
532LogicalResult ConcatOp::inferReturnTypes(
533 MLIRContext *
context, std::optional<Location> loc, ValueRange operands,
534 DictionaryAttr attrs, mlir::PropertyRef properties,
535 mlir::RegionRange regions, SmallVectorImpl<Type> &results) {
537 results.push_back(IntegerType::get(
context, resultWidth));
544ParseResult ConcatOp::parse(OpAsmParser &parser, OperationState &result) {
545 SmallVector<OpAsmParser::UnresolvedOperand, 4> operands;
546 SmallVector<Type, 4> types;
548 llvm::SMLoc allOperandLoc = parser.getCurrentLocation();
551 if (parser.parseOperandList(operands) ||
552 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon())
557 auto parseResult = parser.parseOptionalType(parsedType);
558 if (parseResult.has_value()) {
559 if (failed(parseResult.value()))
561 types.push_back(parsedType);
562 while (succeeded(parser.parseOptionalComma())) {
563 if (parser.parseType(parsedType))
565 types.push_back(parsedType);
569 if (parser.resolveOperands(operands, types, allOperandLoc, result.operands))
572 SmallVector<Type, 1> inferredTypes;
573 if (failed(ConcatOp::inferReturnTypes(
574 parser.getContext(), result.location, result.operands,
575 result.attributes.getDictionary(parser.getContext()),
576 result.getRawProperties(), {}, inferredTypes)))
579 result.addTypes(inferredTypes);
583void ConcatOp::print(OpAsmPrinter &p) {
585 p.printOperands(getOperands());
586 p.printOptionalAttrDict((*this)->getAttrs());
588 llvm::interleaveComma(getOperandTypes(), p);
597OpFoldResult comb::ReverseOp::fold(FoldAdaptor adaptor) {
599 auto cstInput = llvm::dyn_cast_or_null<mlir::IntegerAttr>(adaptor.getInput());
603 APInt val = cstInput.getValue();
604 APInt reversedVal = val.reverseBits();
606 return mlir::IntegerAttr::get(getType(), reversedVal);
613 LogicalResult matchAndRewrite(comb::ReverseOp op,
614 PatternRewriter &rewriter)
const override {
615 auto inputOp = op.getInput().getDefiningOp<comb::ReverseOp>();
619 rewriter.replaceOp(op, inputOp.getInput());
625void comb::ReverseOp::getCanonicalizationPatterns(RewritePatternSet &results,
627 results.add<ReverseOfReverse>(
context);
634LogicalResult ExtractOp::verify() {
635 unsigned srcWidth = cast<IntegerType>(getInput().getType()).getWidth();
636 unsigned dstWidth = cast<IntegerType>(getType()).getWidth();
638 bool checkAddWillOverflow =
639 getLowBit() > std::numeric_limits<
decltype(dstWidth)>::max() - dstWidth;
647 if (checkAddWillOverflow || getLowBit() + dstWidth > srcWidth)
648 return emitOpError(
"from bit too large for input"), failure();
653LogicalResult TruthTableOp::verify() {
654 size_t numInputs = getInputs().size();
655 if (numInputs >=
sizeof(
size_t) * 8)
656 return emitOpError(
"Truth tables support a maximum of ")
657 <<
sizeof(size_t) * 8 - 1 <<
" inputs on your platform";
659 auto table = getLookupTable();
660 if (table.size() != (1ull << numInputs))
661 return emitOpError(
"Expected lookup table of 2^n length");
670#define GET_OP_CLASSES
671#include "circt/Dialect/Comb/Comb.cpp.inc"
assert(baseType &&"element must be base type")
static size_t getTotalWidth(ArrayRef< Value > operands)
static LogicalResult verifyUTBinOp(Operation *op)
static llvm::LogicalResult convertDivModUByPowerOfTwo(PatternRewriter &rewriter, Operation *op, Value lhs, Value rhs, bool isDiv)
static std::unique_ptr< Context > context
Value createOrFoldNot(OpBuilder &builder, Location loc, Value value, bool twoState=false)
Create a `‘Not’' gate on a value.
Value createInject(OpBuilder &builder, Location loc, Value value, unsigned offset, Value replacement)
Replace a range of bits in an integer and return the updated integer value.
Value createZExt(OpBuilder &builder, Location loc, Value value, unsigned targetWidth)
Create the ops to zero-extend a value to an integer of equal or larger type.
Value createOrFoldSExt(OpBuilder &builder, Location loc, Value value, Type destTy)
Create a sign extension operation from a value of integer type to an equal or larger integer type.
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.
void replaceOpAndCopyNamehint(PatternRewriter &rewriter, Operation *op, Value newValue)
A wrapper of PatternRewriter::replaceOp to propagate "sv.namehint" attribute.