401 PatternRewriter &rewriter)
const override {
403 auto isSquarer = op.getLhs() == op.getRhs();
404 if (comb::shouldUseBoothEncoding(op.getLhs(), op.getRhs()) && !isSquarer)
407 auto inputWidth = op.getLhs().getType().getIntOrFloatBitWidth();
410 if (!matchPattern(op.getLhs(), comb::m_SextBy(m_Any(&lhsReplBits))) ||
411 !matchPattern(op.getRhs(), comb::m_SextBy(m_Any(&rhsReplBits))))
415 inputWidth - lhsReplBits.getType().getIntOrFloatBitWidth();
417 inputWidth - rhsReplBits.getType().getIntOrFloatBitWidth();
419 size_t maxRows = std::max(lhsWidth, rhsWidth) - 1;
423 if (lhsWidth != rhsWidth || lhsWidth <= 1 || rhsWidth <= 1)
427 if (maxRows >= op.getNumResults())
431 auto lhsBaseWidth = lhsWidth - 1;
432 auto rhsBaseWidth = rhsWidth - 1;
434 op.getLhs(), lhsBaseWidth, 1);
436 op.getRhs(), rhsBaseWidth, 1);
444 comb::createZExt(rewriter, op.getLoc(), lhsBase, inputWidth);
446 comb::createZExt(rewriter, op.getLoc(), rhsBase, inputWidth);
447 auto newPP = datapath::PartialProductOp::create(
448 rewriter, op.getLoc(), ValueRange{lhsBaseZext, rhsBaseZext}, maxRows);
457 auto lhsSignReplicate = comb::ReplicateOp::create(rewriter, op.getLoc(),
458 lhsSignBit, rhsBaseWidth);
460 comb::AndOp::create(rewriter, op.getLoc(), lhsSignReplicate, rhsBase);
461 auto lhsSignCorrection =
462 comb::createOrFoldNot(rewriter, op.getLoc(), lhsSignAndRhs,
true);
465 auto alignLhsSignCorrection =
zeroPad(
466 rewriter, op.getLoc(), lhsSignCorrection, inputWidth, lhsBaseWidth);
469 auto rhsSignReplicate = comb::ReplicateOp::create(rewriter, op.getLoc(),
470 rhsSignBit, lhsBaseWidth);
472 comb::AndOp::create(rewriter, op.getLoc(), rhsSignReplicate, lhsBase);
473 auto rhsSignCorrection =
474 comb::createOrFoldNot(rewriter, op.getLoc(), rhsSignAndLhs,
true);
477 auto alignRhsSignCorrection =
zeroPad(
478 rewriter, op.getLoc(), rhsSignCorrection, inputWidth, rhsBaseWidth);
483 comb::AndOp::create(rewriter, op.getLoc(), lhsSignBit, rhsSignBit);
485 auto alignSignAndZext =
zeroPad(rewriter, op.getLoc(), signAnd, inputWidth,
486 lhsBaseWidth + rhsBaseWidth);
490 auto ones = APInt::getAllOnes(inputWidth);
491 auto lowerLhs = APInt::getOneBitSet(inputWidth, lhsBaseWidth);
492 auto lowerRhs = APInt::getOneBitSet(inputWidth, rhsBaseWidth);
493 auto msbCorrection = ones << (lhsBaseWidth + rhsBaseWidth);
494 auto correction = lowerLhs + lowerRhs + 2 * msbCorrection;
496 auto constantCorrection =
500 APInt::getZero(inputWidth));
502 SmallVector<Value> newResults(newPP.getResults().begin(),
503 newPP.getResults().end());
506 newResults.push_back(alignLhsSignCorrection);
508 newResults.push_back(alignRhsSignCorrection);
510 newResults.push_back(alignSignAndZext);
512 newResults.push_back(constantCorrection);
514 newResults.append(op.getNumResults() - newResults.size(), zero);
516 rewriter.replaceOp(op, newResults);