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);
186 auto rowWidth = width;
187 auto knownBitsA = comb::computeKnownBits(a);
188 if (!knownBitsA.Zero.isZero()) {
189 if (knownBitsA.Zero.countLeadingOnes() > 1) {
190 rowWidth -= knownBitsA.Zero.countLeadingOnes();
195 SmallVector<Value> partialProducts;
196 partialProducts.reserve(width);
199 assert(op.getNumResults() <= width &&
200 "Cannot return more results than the operator width");
202 for (
unsigned i = 0; i < op.getNumResults(); ++i) {
204 rewriter.createOrFold<comb::ReplicateOp>(loc, bBits[i], rowWidth);
205 auto ppRow = rewriter.createOrFold<
comb::AndOp>(loc, repl, a);
206 if (rowWidth < width) {
207 auto padding = width - rowWidth;
210 loc, ValueRange{
zeroPad, ppRow});
214 partialProducts.push_back(ppRow);
219 comb::ConcatOp::create(rewriter, loc, ValueRange{ppRow, shiftBy});
221 loc, ppAlign, 0, width);
222 partialProducts.push_back(ppAlignTrunc);
225 rewriter.replaceOp(op, partialProducts);
229 static LogicalResult lowerSqrAndArray(PatternRewriter &rewriter, Value a,
230 PartialProductOp op,
unsigned width) {
232 Location loc = op.getLoc();
233 SmallVector<Value> aBits =
extractBits(rewriter, a);
235 SmallVector<Value> partialProducts;
236 partialProducts.reserve(width);
240 assert(op.getNumResults() <= width &&
241 "Cannot return more results than the operator width");
243 for (
unsigned i = 0; i < op.getNumResults(); ++i) {
244 SmallVector<Value> row;
247 if (2 * i >= width) {
250 partialProducts.push_back(zeroWidth);
256 row.push_back(shiftBy);
258 row.push_back(aBits[i]);
261 unsigned rowWidth = 2 * i + 1;
262 if (rowWidth < width) {
263 row.push_back(zeroFalse);
267 for (
unsigned j = i + 1; j < width; ++j) {
269 if (rowWidth == width)
275 if (j >= op.getNumResults()) {
276 row.push_back(zeroFalse);
281 rewriter.createOrFold<
comb::AndOp>(loc, aBits[i], aBits[j]);
282 row.push_back(ppBit);
284 std::reverse(row.begin(), row.end());
285 auto ppRow = comb::ConcatOp::create(rewriter, loc, row);
286 partialProducts.push_back(ppRow);
289 rewriter.replaceOp(op, partialProducts);
293 static LogicalResult lowerBoothArray(PatternRewriter &rewriter, Value a,
294 Value b, PartialProductOp op,
297 Location loc = op.getLoc();
300 auto [aSigned, aBase] =
getBaseOfExt(rewriter, loc, op.getLhs());
301 auto [bSigned, bBase] =
getBaseOfExt(rewriter, loc, op.getRhs());
303 auto aBaseWidth = aBase.getType().getIntOrFloatBitWidth();
304 auto bBaseWidth = bBase.getType().getIntOrFloatBitWidth();
308 auto rowWidth = width;
309 if (aBaseWidth < width) {
311 rowWidth = aBaseWidth + 1;
317 rewriter.createOrFold<
comb::ConcatOp>(loc, ValueRange{a, zeroFalse});
319 loc, twoAPre, 0, rowWidth);
323 SmallVector<Value> bBits =
extractBits(rewriter, b);
325 bBits.push_back(zeroFalse);
326 bBits.push_back(zeroFalse);
331 bBits.resize(bBaseWidth + 2);
335 bBits.resize(bBaseWidth + 1);
337 SmallVector<Value> partialProducts;
338 partialProducts.reserve(op.getNumResults());
345 SmallVector<Value> encNegs;
349 for (
unsigned i = 0; i + 1 < bBits.size(); i += 2) {
351 Value bim1 = (i == 0) ? zeroFalse : bBits[i - 1];
353 Value bip1 = bBits[i + 1];
357 encNegs.push_back(encNeg);
359 Value encOne = rewriter.createOrFold<
comb::XorOp>(loc, bi, bim1,
true);
362 Value biInv = rewriter.createOrFold<
comb::XorOp>(loc, bi, constOne,
true);
364 rewriter.createOrFold<
comb::XorOp>(loc, bip1, constOne,
true);
366 rewriter.createOrFold<
comb::XorOp>(loc, bim1, constOne,
true);
368 Value andLeft = rewriter.createOrFold<
comb::AndOp>(
369 loc, ValueRange{bip1Inv, bi, bim1},
true);
370 Value andRight = rewriter.createOrFold<
comb::AndOp>(
371 loc, ValueRange{bip1, biInv, bim1Inv},
true);
373 rewriter.createOrFold<
comb::OrOp>(loc, andLeft, andRight,
true);
376 rewriter.createOrFold<comb::ReplicateOp>(loc, encNeg, rowWidth);
378 rewriter.createOrFold<comb::ReplicateOp>(loc, encOne, rowWidth);
380 rewriter.createOrFold<comb::ReplicateOp>(loc, encTwo, rowWidth);
383 Value selTwoA = rewriter.createOrFold<
comb::AndOp>(loc, encTwoRepl, twoA);
384 Value selOneA = rewriter.createOrFold<
comb::AndOp>(loc, encOneRepl, a);
386 rewriter.createOrFold<
comb::OrOp>(loc, selTwoA, selOneA,
true);
390 rewriter.createOrFold<
comb::XorOp>(loc, magA, encNegRepl,
true);
394 partialProducts.push_back(ppRow);
401 loc, ValueRange{ppRow, zeroFalse, encNegPrev});
402 partialProducts.push_back(withSignCorrection);
411 loc, ValueRange{ppRow, zeroFalse, encNegPrev, shiftBy});
412 partialProducts.push_back(withSignCorrection);
415 if (partialProducts.size() == op.getNumResults())
422 auto numPP = partialProducts.size();
425 Value finalSignCorrection = rewriter.createOrFold<
comb::ConcatOp>(
426 loc, ValueRange{zeroFalse, encNegPrev, shiftByFinal});
427 partialProducts.push_back(finalSignCorrection);
428 encNegs.push_back(zeroFalse);
438 for (
unsigned i = 0; i < partialProducts.size(); ++i) {
439 auto ppRow = partialProducts[i];
441 auto ppWidth = ppRow.getType().getIntOrFloatBitWidth();
442 if (ppWidth < width) {
443 auto padding = width - ppWidth;
444 auto encNeg = encNegs[i];
451 rewriter.createOrFold<comb::ReplicateOp>(loc, encNeg, padding);
453 loc, ValueRange{encNegPad, ppRow});
457 ppWidth = ppRow.getType().getIntOrFloatBitWidth();
458 if (ppWidth > width) {
461 partialProducts[i] = ppRow;
462 assert(partialProducts[i].getType().getIntOrFloatBitWidth() == width &&
463 "Expected sign-extended partial product to be full width");
468 while (partialProducts.size() < op.getNumResults())
469 partialProducts.push_back(zeroWidth);
471 assert(partialProducts.size() == op.getNumResults() &&
472 "Expected number of booth partial products to match results");
474 rewriter.replaceOp(op, partialProducts);
479struct DatapathPosPartialProductOpConversion
483 DatapathPosPartialProductOpConversion(MLIRContext *
context,
bool forceBooth)
485 forceBooth(forceBooth){};
487 const bool forceBooth;
489 LogicalResult matchAndRewrite(PosPartialProductOp op,
490 PatternRewriter &rewriter)
const override {
492 Value a = op.getAddend0();
493 Value b = op.getAddend1();
494 Value c = op.getMultiplicand();
495 unsigned width = a.getType().getIntOrFloatBitWidth();
499 rewriter.replaceOpWithNewOp<
hw::ConstantOp>(op, op.getType(0), 0);
504 return lowerAndArray(rewriter, a, b, c, op, width);
508 static LogicalResult lowerAndArray(PatternRewriter &rewriter, Value a,
509 Value b, Value c, PosPartialProductOp op,
512 Location loc = op.getLoc();
515 auto carry = rewriter.createOrFold<
comb::AndOp>(loc, a, b);
516 auto save = rewriter.createOrFold<
comb::XorOp>(loc, a, b);
518 SmallVector<Value> carryBits =
extractBits(rewriter, carry);
519 SmallVector<Value> saveBits =
extractBits(rewriter, save);
522 auto rowWidth = width;
524 auto cBaseWidth = cBase.getType().getIntOrFloatBitWidth();
526 if (cBaseWidth < width && !cSigned) {
528 rowWidth = cBaseWidth + 1;
535 comb::ConcatOp::create(rewriter, loc, ValueRange{c, zeroFalse});
541 SmallVector<Value> partialProducts;
542 partialProducts.reserve(width);
544 assert(op.getNumResults() <= width &&
545 "Cannot return more results than the operator width");
547 for (
unsigned i = 0; i < op.getNumResults(); ++i) {
549 rewriter.createOrFold<comb::ReplicateOp>(loc, saveBits[i], rowWidth);
551 rewriter.createOrFold<comb::ReplicateOp>(loc, carryBits[i], rowWidth);
553 auto ppRowSave = rewriter.createOrFold<
comb::AndOp>(loc, replSave, c);
555 rewriter.createOrFold<
comb::AndOp>(loc, replCarry, twoC);
557 rewriter.createOrFold<
comb::OrOp>(loc, ppRowSave, ppRowCarry);
558 auto ppAlign = ppRow;
562 comb::ConcatOp::create(rewriter, loc, ValueRange{ppRow, shiftBy});
566 if (rowWidth + i > width) {
569 partialProducts.push_back(ppAlignTrunc);
573 if (rowWidth + i < width) {
574 auto extPPAlign = comb::createZExt(rewriter, loc, ppAlign, width);
575 partialProducts.push_back(extPPAlign);
579 partialProducts.push_back(ppAlign);
582 rewriter.replaceOp(op, partialProducts);
594struct ConvertDatapathToCombPass
595 :
public impl::ConvertDatapathToCombBase<ConvertDatapathToCombPass> {
596 void runOnOperation()
override;
597 using ConvertDatapathToCombBase<
598 ConvertDatapathToCombPass>::ConvertDatapathToCombBase;
603 Operation *op, RewritePatternSet &&
patterns,
607 mlir::GreedyRewriteConfig config;
612 config.setMaxIterations(2).setListener(analysis).setUseTopDownTraversal(
true);
615 if (failed(mlir::applyPatternsGreedily(op, std::move(
patterns), config)))
621void ConvertDatapathToCombPass::runOnOperation() {
622 RewritePatternSet
patterns(&getContext());
624 patterns.add<DatapathPartialProductOpConversion,
625 DatapathPosPartialProductOpConversion>(
patterns.getContext(),
629 analysis = &getAnalysis<synth::IncrementalLongestPathAnalysis>();
631 if (lowerCompressToAdd)
636 patterns.add<DatapathCompressOpConversion>(
patterns.getContext(), analysis);
639 getOperation(), std::move(
patterns), analysis)))
640 return signalPassFailure();
645 auto result = getOperation()->walk([&](Operation *op) {
646 if (llvm::isa<datapath::CompressOp>(op) && !lowerCompress &&
648 return WalkResult::advance();
649 if (llvm::isa_and_nonnull<datapath::DatapathDialect>(op->getDialect())) {
650 op->emitError(
"Datapath operation not converted: ") << *op;
651 return WalkResult::interrupt();
653 return WalkResult::advance();
655 if (result.wasInterrupted())
656 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.