14#include "mlir/Analysis/TopologicalSortUtils.h"
15#include "mlir/Dialect/Func/IR/FuncOps.h"
16#include "mlir/IR/PatternMatch.h"
17#include "mlir/Pass/Pass.h"
18#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
19#include "llvm/Support/Debug.h"
20#include "llvm/Support/KnownBits.h"
23#define DEBUG_TYPE "datapath-to-comb"
26#define GEN_PASS_DEF_CONVERTDATAPATHTOCOMB
27#include "circt/Conversion/Passes.h.inc"
31using namespace datapath;
34static SmallVector<Value>
extractBits(OpBuilder &builder, Value val) {
35 SmallVector<Value> bits;
36 comb::extractBits(builder, val, bits);
42static std::pair<bool, Value>
getBaseOfExt(PatternRewriter &rewriter,
43 Location loc, Value val) {
47 if (matchPattern(val, comb::m_ZextBy(mlir::matchers::m_Any(&replBits)))) {
48 auto baseWidth = val.getType().getIntOrFloatBitWidth() -
49 replBits.getType().getIntOrFloatBitWidth();
52 return {
false, valBase};
56 if (matchPattern(val, comb::m_SextBy(mlir::matchers::m_Any(&replBits)))) {
57 auto baseWidth = val.getType().getIntOrFloatBitWidth() -
58 replBits.getType().getIntOrFloatBitWidth();
61 return {
true, valBase};
79 matchAndRewrite(CompressOp op,
80 mlir::PatternRewriter &rewriter)
const override {
81 Location loc = op.getLoc();
82 auto inputs = op.getOperands();
83 unsigned width = inputs[0].getType().getIntOrFloatBitWidth();
85 auto addOp = comb::AddOp::create(rewriter, loc, inputs,
true);
88 SmallVector<Value> results(op.getNumResults() - 1, zeroOp);
89 results.push_back(addOp);
90 rewriter.replaceOp(op, results);
97 DatapathCompressOpConversion(MLIRContext *
context,
102 matchAndRewrite(CompressOp op,
103 mlir::PatternRewriter &rewriter)
const override {
104 Location loc = op.getLoc();
105 auto inputs = op.getOperands();
107 SmallVector<SmallVector<Value>> addends;
108 for (
auto input : inputs) {
114 auto width = inputs[0].getType().getIntOrFloatBitWidth();
115 auto targetAddends = op.getNumResults();
120 if (failed(comp.withInputDelays(
121 [&](Value v) { return analysis->getMaxDelay(v, 0); })))
125 rewriter.replaceOp(op, comp.compressToHeight(rewriter, targetAddends));
133struct DatapathPartialProductOpConversion :
OpRewritePattern<PartialProductOp> {
136 DatapathPartialProductOpConversion(MLIRContext *
context,
bool forceBooth)
139 const bool forceBooth;
141 LogicalResult matchAndRewrite(PartialProductOp op,
142 PatternRewriter &rewriter)
const override {
144 Value a = op.getLhs();
145 Value b = op.getRhs();
146 unsigned width = a.getType().getIntOrFloatBitWidth();
150 rewriter.replaceOpWithNewOp<
hw::ConstantOp>(op, op.getType(0), 0);
167 return lowerSqrAndArray(rewriter, a, op, width);
171 if (comb::shouldUseBoothEncoding(a, b) || forceBooth)
172 return lowerBoothArray(rewriter, a, b, op, width);
174 return lowerAndArray(rewriter, a, b, op, width);
178 static LogicalResult lowerAndArray(PatternRewriter &rewriter, Value a,
179 Value b, PartialProductOp op,
182 Location loc = op.getLoc();
184 SmallVector<Value> bBits =
extractBits(rewriter, b);
185 auto knownBitsB = comb::computeKnownBits(b);
187 auto rowWidth = width;
188 auto knownBitsA = comb::computeKnownBits(a);
189 if (!knownBitsA.Zero.isZero()) {
190 if (knownBitsA.Zero.countLeadingOnes() > 1) {
191 rowWidth -= knownBitsA.Zero.countLeadingOnes();
196 SmallVector<Value> partialProducts;
197 partialProducts.reserve(width);
200 assert(op.getNumResults() <= width &&
201 "Cannot return more results than the operator width");
203 for (
unsigned i = 0; i < op.getNumResults(); ++i) {
208 if (knownBitsB.Zero[i]) {
209 partialProducts.push_back(
215 if (knownBitsB.One[i]) {
219 rewriter.createOrFold<comb::ReplicateOp>(loc, bBits[i], rowWidth);
220 ppRow = rewriter.createOrFold<
comb::AndOp>(loc, repl, a);
222 if (rowWidth < width) {
223 auto padding = width - rowWidth;
226 loc, ValueRange{
zeroPad, ppRow});
230 partialProducts.push_back(ppRow);
235 comb::ConcatOp::create(rewriter, loc, ValueRange{ppRow, shiftBy});
237 loc, ppAlign, 0, width);
238 partialProducts.push_back(ppAlignTrunc);
241 rewriter.replaceOp(op, partialProducts);
245 static LogicalResult lowerSqrAndArray(PatternRewriter &rewriter, Value a,
246 PartialProductOp op,
unsigned width) {
248 Location loc = op.getLoc();
249 SmallVector<Value> aBits =
extractBits(rewriter, a);
251 SmallVector<Value> partialProducts;
252 partialProducts.reserve(width);
256 assert(op.getNumResults() <= width &&
257 "Cannot return more results than the operator width");
259 for (
unsigned i = 0; i < op.getNumResults(); ++i) {
260 SmallVector<Value> row;
263 if (2 * i >= width) {
266 partialProducts.push_back(zeroWidth);
272 row.push_back(shiftBy);
274 row.push_back(aBits[i]);
277 unsigned rowWidth = 2 * i + 1;
278 if (rowWidth < width) {
279 row.push_back(zeroFalse);
283 for (
unsigned j = i + 1; j < width; ++j) {
285 if (rowWidth == width)
291 if (j >= op.getNumResults()) {
292 row.push_back(zeroFalse);
297 rewriter.createOrFold<
comb::AndOp>(loc, aBits[i], aBits[j]);
298 row.push_back(ppBit);
300 std::reverse(row.begin(), row.end());
301 auto ppRow = comb::ConcatOp::create(rewriter, loc, row);
302 partialProducts.push_back(ppRow);
305 rewriter.replaceOp(op, partialProducts);
309 static LogicalResult lowerBoothArray(PatternRewriter &rewriter, Value a,
310 Value b, PartialProductOp op,
313 Location loc = op.getLoc();
316 auto [aSigned, aBase] =
getBaseOfExt(rewriter, loc, op.getLhs());
317 auto [bSigned, bBase] =
getBaseOfExt(rewriter, loc, op.getRhs());
319 auto aBaseWidth = aBase.getType().getIntOrFloatBitWidth();
320 auto bBaseWidth = bBase.getType().getIntOrFloatBitWidth();
324 auto rowWidth = width;
325 if (aBaseWidth < width) {
327 rowWidth = aBaseWidth + 1;
333 rewriter.createOrFold<
comb::ConcatOp>(loc, ValueRange{a, zeroFalse});
335 loc, twoAPre, 0, rowWidth);
339 SmallVector<Value> bBits =
extractBits(rewriter, b);
341 bBits.append(2, zeroFalse);
346 bBits.resize(bBaseWidth + 2);
350 bBits.resize(bBaseWidth + 1);
352 SmallVector<Value> partialProducts;
353 partialProducts.reserve(op.getNumResults());
360 SmallVector<Value> encNegs;
364 for (
unsigned i = 0; i + 1 < bBits.size(); i += 2) {
366 Value bim1 = (i == 0) ? zeroFalse : bBits[i - 1];
368 Value bip1 = bBits[i + 1];
372 encNegs.push_back(encNeg);
374 Value encOne = rewriter.createOrFold<
comb::XorOp>(loc, bi, bim1,
true);
377 Value biInv = rewriter.createOrFold<
comb::XorOp>(loc, bi, constOne,
true);
379 rewriter.createOrFold<
comb::XorOp>(loc, bip1, constOne,
true);
381 rewriter.createOrFold<
comb::XorOp>(loc, bim1, constOne,
true);
383 Value andLeft = rewriter.createOrFold<
comb::AndOp>(
384 loc, ValueRange{bip1Inv, bi, bim1},
true);
385 Value andRight = rewriter.createOrFold<
comb::AndOp>(
386 loc, ValueRange{bip1, biInv, bim1Inv},
true);
388 rewriter.createOrFold<
comb::OrOp>(loc, andLeft, andRight,
true);
391 rewriter.createOrFold<comb::ReplicateOp>(loc, encNeg, rowWidth);
393 rewriter.createOrFold<comb::ReplicateOp>(loc, encOne, rowWidth);
395 rewriter.createOrFold<comb::ReplicateOp>(loc, encTwo, rowWidth);
398 Value selTwoA = rewriter.createOrFold<
comb::AndOp>(loc, encTwoRepl, twoA);
399 Value selOneA = rewriter.createOrFold<
comb::AndOp>(loc, encOneRepl, a);
401 rewriter.createOrFold<
comb::OrOp>(loc, selTwoA, selOneA,
true);
405 rewriter.createOrFold<
comb::XorOp>(loc, magA, encNegRepl,
true);
409 partialProducts.push_back(ppRow);
416 loc, ValueRange{ppRow, zeroFalse, encNegPrev});
417 partialProducts.push_back(withSignCorrection);
426 loc, ValueRange{ppRow, zeroFalse, encNegPrev, shiftBy});
427 partialProducts.push_back(withSignCorrection);
430 if (partialProducts.size() == op.getNumResults())
437 auto numPP = partialProducts.size();
440 Value finalSignCorrection = rewriter.createOrFold<
comb::ConcatOp>(
441 loc, ValueRange{zeroFalse, encNegPrev, shiftByFinal});
442 partialProducts.push_back(finalSignCorrection);
443 encNegs.push_back(zeroFalse);
453 for (
unsigned i = 0; i < partialProducts.size(); ++i) {
454 auto ppRow = partialProducts[i];
456 auto ppWidth = ppRow.getType().getIntOrFloatBitWidth();
457 if (ppWidth < width) {
458 auto padding = width - ppWidth;
459 auto encNeg = encNegs[i];
466 rewriter.createOrFold<comb::ReplicateOp>(loc, encNeg, padding);
468 loc, ValueRange{encNegPad, ppRow});
472 ppWidth = ppRow.getType().getIntOrFloatBitWidth();
473 if (ppWidth > width) {
476 partialProducts[i] = ppRow;
477 assert(partialProducts[i].getType().getIntOrFloatBitWidth() == width &&
478 "Expected sign-extended partial product to be full width");
483 while (partialProducts.size() < op.getNumResults())
484 partialProducts.push_back(zeroWidth);
486 assert(partialProducts.size() == op.getNumResults() &&
487 "Expected number of booth partial products to match results");
489 rewriter.replaceOp(op, partialProducts);
494struct DatapathPosPartialProductOpConversion
498 DatapathPosPartialProductOpConversion(MLIRContext *
context,
bool forceBooth)
500 forceBooth(forceBooth){};
502 const bool forceBooth;
504 LogicalResult matchAndRewrite(PosPartialProductOp op,
505 PatternRewriter &rewriter)
const override {
507 Value a = op.getAddend0();
508 Value b = op.getAddend1();
509 Value c = op.getMultiplicand();
510 unsigned width = a.getType().getIntOrFloatBitWidth();
514 rewriter.replaceOpWithNewOp<
hw::ConstantOp>(op, op.getType(0), 0);
521 return lowerAndArray(rewriter, a, b, c, op, width);
525 if (comb::shouldUseBoothEncoding(a, b) || forceBooth)
526 return lowerBoothArray(rewriter, a, b, c, op, width);
527 return lowerAndArray(rewriter, a, b, c, op, width);
531 static LogicalResult lowerBoothArray(PatternRewriter &rewriter, Value a,
532 Value b, Value c, PosPartialProductOp op,
534 Location loc = op.getLoc();
551 unsigned cBaseWidth = cBase.getType().getIntOrFloatBitWidth();
552 unsigned rowWidth = width;
553 if (cBaseWidth < width) {
555 rowWidth = cBaseWidth + 1;
568 auto aBaseWidth = aBase.getType().getIntOrFloatBitWidth();
569 auto bBaseWidth = bBase.getType().getIntOrFloatBitWidth();
571 auto encodeSigned = aSigned && bSigned;
572 auto encodeBaseWidth = std::max(aBaseWidth, bBaseWidth);
574 SmallVector<Value> aBits =
extractBits(rewriter, a);
575 SmallVector<Value> bBits =
extractBits(rewriter, b);
577 aBits.resize(encodeBaseWidth);
578 bBits.resize(encodeBaseWidth);
582 aBits.append(3, zero);
583 bBits.append(3, zero);
588 aBits.append(2, aBits.back());
589 bBits.append(2, bBits.back());
591 SmallVector<Value> partialProducts;
592 SmallVector<Value> encNegs;
593 partialProducts.reserve(op.getNumResults());
594 encNegs.reserve(aBits.size());
596 Value recoderCarry = zero;
598 for (
unsigned i = 0; i + 1 < aBits.size(); i += 2) {
600 Value aim1 = (i == 0) ? zero : aBits[i - 1];
601 Value bim1 = (i == 0) ? zero : bBits[i - 1];
604 Value aip1 = aBits[i + 1];
605 Value bip1 = bBits[i + 1];
616 Value aAndB = rewriter.createOrFold<
comb::AndOp>(loc, ai, bi,
true);
617 Value aAndPrevB = rewriter.createOrFold<
comb::AndOp>(loc, ai, bim1,
true);
618 Value bAndPrevB = rewriter.createOrFold<
comb::AndOp>(loc, bi, bim1,
true);
619 Value majority = rewriter.createOrFold<
comb::OrOp>(
620 loc, ValueRange{aAndB, aAndPrevB, bAndPrevB},
true);
622 Value nextXor = rewriter.createOrFold<
comb::XorOp>(loc, aip1, bip1,
true);
623 Value aXorB = rewriter.createOrFold<
comb::XorOp>(loc, ai, bi,
true);
624 Value aXorPrevA = rewriter.createOrFold<
comb::XorOp>(loc, ai, aim1,
true);
625 Value prevOr = rewriter.createOrFold<
comb::OrOp>(loc, aim1, bim1,
true);
629 Value y1 = rewriter.createOrFold<
comb::XorOp>(loc, aXorB, prevOr,
true);
630 Value y2 = rewriter.createOrFold<
comb::OrOp>(loc, aXorB, aXorPrevA,
true);
632 Value z1 = rewriter.createOrFold<
comb::XorOp>(loc, nextXor, y2,
true);
637 rewriter.createOrFold<
comb::XorOp>(loc, majority, nextXor,
true);
638 Value invNextXor = comb::createOrFoldNot(rewriter, loc, nextXor);
641 Value recoderCarryNext =
642 rewriter.createOrFold<
comb::AndOp>(loc, invNextXor, majority,
true);
646 loc, ValueRange{y1, recoderCarry},
true);
647 Value encOneInv = comb::createOrFoldNot(rewriter, loc, encOne);
650 loc, ValueRange{z1, encOneInv},
true);
654 rewriter.createOrFold<comb::ReplicateOp>(loc, encNeg, rowWidth);
656 rewriter.createOrFold<comb::ReplicateOp>(loc, encOne, rowWidth);
658 rewriter.createOrFold<comb::ReplicateOp>(loc, encTwo, rowWidth);
659 Value selTwoC = rewriter.createOrFold<
comb::AndOp>(loc, encTwoRepl, twoC);
660 Value selOneC = rewriter.createOrFold<
comb::AndOp>(loc, encOneRepl, c);
661 Value magnitude = rewriter.createOrFold<
comb::OrOp>(
662 loc, ValueRange{selTwoC, selOneC},
true);
664 rewriter.createOrFold<
comb::XorOp>(loc, magnitude, encNegRepl,
true);
666 encNegs.push_back(encNeg);
668 partialProducts.push_back(ppRow);
671 loc, ValueRange{ppRow, zero, encNegPrev}));
675 loc, ValueRange{ppRow, zero, encNegPrev, shift}));
678 recoderCarry = recoderCarryNext;
680 if (partialProducts.size() == op.getNumResults())
688 auto numPP = partialProducts.size();
691 Value finalSignCorrection = rewriter.createOrFold<
comb::ConcatOp>(
692 loc, ValueRange{zero, encNegPrev, shiftByFinal});
693 partialProducts.push_back(finalSignCorrection);
694 encNegs.push_back(zero);
698 for (
auto [index, ppRow] :
llvm::enumerate(partialProducts)) {
699 unsigned ppWidth = ppRow.getType().getIntOrFloatBitWidth();
700 if (ppWidth < width) {
701 Value sign = encNegs[index];
705 Value padding = rewriter.createOrFold<comb::ReplicateOp>(
706 loc, sign, width - ppWidth);
708 loc, ValueRange{padding, ppRow});
710 if (ppRow.getType().getIntOrFloatBitWidth() > width)
712 partialProducts[index] = ppRow;
716 partialProducts.resize(op.getNumResults(), zeroWidth);
717 rewriter.replaceOp(op, partialProducts);
721 static LogicalResult lowerAndArray(PatternRewriter &rewriter, Value a,
722 Value b, Value c, PosPartialProductOp op,
725 Location loc = op.getLoc();
728 auto carry = rewriter.createOrFold<
comb::AndOp>(loc, a, b);
729 auto save = rewriter.createOrFold<
comb::XorOp>(loc, a, b);
731 SmallVector<Value> carryBits =
extractBits(rewriter, carry);
732 SmallVector<Value> saveBits =
extractBits(rewriter, save);
735 auto rowWidth = width;
737 auto cBaseWidth = cBase.getType().getIntOrFloatBitWidth();
739 if (cBaseWidth < width && !cSigned) {
741 rowWidth = cBaseWidth + 1;
748 comb::ConcatOp::create(rewriter, loc, ValueRange{c, zeroFalse});
754 SmallVector<Value> partialProducts;
755 partialProducts.reserve(width);
757 assert(op.getNumResults() <= width &&
758 "Cannot return more results than the operator width");
760 for (
unsigned i = 0; i < op.getNumResults(); ++i) {
762 rewriter.createOrFold<comb::ReplicateOp>(loc, saveBits[i], rowWidth);
764 rewriter.createOrFold<comb::ReplicateOp>(loc, carryBits[i], rowWidth);
766 auto ppRowSave = rewriter.createOrFold<
comb::AndOp>(loc, replSave, c);
768 rewriter.createOrFold<
comb::AndOp>(loc, replCarry, twoC);
770 rewriter.createOrFold<
comb::OrOp>(loc, ppRowSave, ppRowCarry);
771 auto ppAlign = ppRow;
775 comb::ConcatOp::create(rewriter, loc, ValueRange{ppRow, shiftBy});
779 if (rowWidth + i > width) {
782 partialProducts.push_back(ppAlignTrunc);
786 if (rowWidth + i < width) {
787 auto extPPAlign = comb::createZExt(rewriter, loc, ppAlign, width);
788 partialProducts.push_back(extPPAlign);
792 partialProducts.push_back(ppAlign);
795 rewriter.replaceOp(op, partialProducts);
807struct ConvertDatapathToCombPass
808 :
public impl::ConvertDatapathToCombBase<ConvertDatapathToCombPass> {
809 void runOnOperation()
override;
810 using ConvertDatapathToCombBase<
811 ConvertDatapathToCombPass>::ConvertDatapathToCombBase;
816 Operation *op, RewritePatternSet &&
patterns,
820 mlir::GreedyRewriteConfig config;
825 config.setMaxIterations(2).setListener(analysis).setUseTopDownTraversal(
true);
828 if (failed(mlir::applyPatternsGreedily(op, std::move(
patterns), config)))
834void ConvertDatapathToCombPass::runOnOperation() {
835 RewritePatternSet
patterns(&getContext());
837 patterns.add<DatapathPartialProductOpConversion,
838 DatapathPosPartialProductOpConversion>(
patterns.getContext(),
842 analysis = &getAnalysis<synth::IncrementalLongestPathAnalysis>();
844 if (lowerCompressToAdd)
849 patterns.add<DatapathCompressOpConversion>(
patterns.getContext(), analysis);
852 getOperation(), std::move(
patterns), analysis)))
853 return signalPassFailure();
858 auto result = getOperation()->walk([&](Operation *op) {
859 if (llvm::isa<datapath::CompressOp>(op) && !lowerCompress &&
861 return WalkResult::advance();
862 if (llvm::isa_and_nonnull<datapath::DatapathDialect>(op->getDialect())) {
863 op->emitError(
"Datapath operation not converted: ") << *op;
864 return WalkResult::interrupt();
866 return WalkResult::advance();
868 if (result.wasInterrupted())
869 return signalPassFailure();
assert(baseType &&"element must be base type")
static SmallVector< Value > extractBits(OpBuilder &builder, Value val)
static Value zeroPad(PatternRewriter &rewriter, Location loc, Value input, size_t targetWidth, size_t trailingZeros)
static std::pair< bool, Value > getBaseOfExt(PatternRewriter &rewriter, Location loc, Value val)
static SmallVector< Value > extractBits(OpBuilder &builder, Value val)
static LogicalResult applyPatternsGreedilyWithTimingInfo(Operation *op, RewritePatternSet &&patterns, synth::IncrementalLongestPathAnalysis *analysis)
static std::unique_ptr< Context > context
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.