From e8a51c0abe98711e947bcc36fc44b33ed5f93a53 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Markus=20B=C3=B6ck?= Date: Thu, 26 Mar 2026 11:51:27 +0100 Subject: [PATCH 1/3] [HandshakeOptimizeBitwidths] Fix forward logic for `shrsi` The previous logic for `shrsi` for the forward pass often crashed in edge cases such as the shift amount was larger than the bitwidth. This PR rewrites the forward logic for `shrsi` into a dedicated pattern. In the case of the input of `shrsi` being zero-extended we just optimize it to a `shrui` and reuse the existing optimization logic there. Fixes https://github.com/EPFL-LAP/dynamatic/issues/792 --- .../dynamatic/Dialect/Handshake/Handshake.td | 1 + .../Dialect/Handshake/HandshakeArithOps.td | 3 - .../Handshake/HandshakeCanonicalization.td | 9 ++ lib/Dialect/Handshake/HandshakeOps.cpp | 20 +-- lib/Transforms/HandshakeOptimizeBitwidths.cpp | 153 +++++++++++------- 5 files changed, 111 insertions(+), 75 deletions(-) diff --git a/include/dynamatic/Dialect/Handshake/Handshake.td b/include/dynamatic/Dialect/Handshake/Handshake.td index e094f57c61..69b8d397c8 100644 --- a/include/dynamatic/Dialect/Handshake/Handshake.td +++ b/include/dynamatic/Dialect/Handshake/Handshake.td @@ -41,6 +41,7 @@ def Handshake_Dialect : Dialect { let useDefaultTypePrinterParser = 1; let useDefaultAttributePrinterParser = 1; let usePropertiesForAttributes = 0; + let hasCanonicalizer = 1; } include "dynamatic/Dialect/Handshake/HandshakeAttributes.td" diff --git a/include/dynamatic/Dialect/Handshake/HandshakeArithOps.td b/include/dynamatic/Dialect/Handshake/HandshakeArithOps.td index 9b855154f3..fda3021691 100644 --- a/include/dynamatic/Dialect/Handshake/HandshakeArithOps.td +++ b/include/dynamatic/Dialect/Handshake/HandshakeArithOps.td @@ -310,12 +310,10 @@ def Handshake_MinUIOp : Handshake_Arith_IntBinaryOp<"minui"> { def Handshake_ExtSIOp : Handshake_Arith_IToICastOp<"extsi"> { let summary = "Integer unsigned width extension."; - let hasCanonicalizer = 1; } def Handshake_ExtUIOp : Handshake_Arith_IToICastOp<"extui"> { let summary = "Integer signed width extension."; - let hasCanonicalizer = 1; } def Handshake_MaximumFOp : Handshake_Arith_FloatBinaryOp<"maximumf", [ @@ -410,7 +408,6 @@ def Handshake_SubIOp : Handshake_Arith_IntBinaryOp<"subi"> { def Handshake_TruncIOp : Handshake_Arith_IToICastOp<"trunci"> { let summary = "Integer truncation."; - let hasCanonicalizer = 1; } def Handshake_TruncFOp : Handshake_Arith_FToFCastOp<"truncf", [ diff --git a/lib/Dialect/Handshake/HandshakeCanonicalization.td b/lib/Dialect/Handshake/HandshakeCanonicalization.td index 0fc4846274..125ba1ded0 100644 --- a/lib/Dialect/Handshake/HandshakeCanonicalization.td +++ b/lib/Dialect/Handshake/HandshakeCanonicalization.td @@ -60,4 +60,13 @@ def TruncIExtUIToExtUI : Pat< [(ValueWiderThan $ext, $tr), (ValueWiderThan $tr, $x)] >; +//===----------------------------------------------------------------------===// +// ShrSIOp +//===----------------------------------------------------------------------===// + +def ShrSIOpOfExtUI : Pat< + (Handshake_ShRSIOp (Handshake_ExtUIOp:$ext $_), $shift), + (Handshake_ShRUIOp $ext, $shift) +>; + #endif // DYNAMATIC_DIALECT_HANDSHAKE_HANDSHAKE_CANONICALIZATION_TD diff --git a/lib/Dialect/Handshake/HandshakeOps.cpp b/lib/Dialect/Handshake/HandshakeOps.cpp index 3bfcaae04a..63813703e1 100644 --- a/lib/Dialect/Handshake/HandshakeOps.cpp +++ b/lib/Dialect/Handshake/HandshakeOps.cpp @@ -206,6 +206,11 @@ namespace { #include "lib/Dialect/Handshake/HandshakeCanonicalization.inc" } // namespace +void handshake::HandshakeDialect::getCanonicalizationPatterns( + RewritePatternSet &set) const { + populateWithGenerated(set); +} + //===----------------------------------------------------------------------===// // MergeOp //===----------------------------------------------------------------------===// @@ -1896,11 +1901,6 @@ static OpFoldResult foldExtOp(Op op) { return nullptr; } -void ExtSIOp::getCanonicalizationPatterns(RewritePatternSet &results, - MLIRContext *context) { - results.add(context); -} - OpFoldResult ExtSIOp::fold(FoldAdaptor adaptor) { return foldExtOp(*this); } /// Extension operations can only extend to a channel with a wider data type and @@ -1927,22 +1927,12 @@ LogicalResult ExtSIOp::verify() { return verifyExtOp(*this); } OpFoldResult ExtUIOp::fold(FoldAdaptor adaptor) { return foldExtOp(*this); } -void ExtUIOp::getCanonicalizationPatterns(RewritePatternSet &results, - MLIRContext *context) { - results.add(context); -} - LogicalResult ExtUIOp::verify() { return verifyExtOp(*this); } //===----------------------------------------------------------------------===// // TruncIOp //===----------------------------------------------------------------------===// -void TruncIOp::getCanonicalizationPatterns(RewritePatternSet &results, - MLIRContext *context) { - results.add(context); -} - OpFoldResult TruncIOp::fold(FoldAdaptor adaptor) { if (auto defTruncOp = getIn().getDefiningOp()) { // Bypass the preceeding truncation operation diff --git a/lib/Transforms/HandshakeOptimizeBitwidths.cpp b/lib/Transforms/HandshakeOptimizeBitwidths.cpp index 4eec5c8c2c..1900c72af2 100644 --- a/lib/Transforms/HandshakeOptimizeBitwidths.cpp +++ b/lib/Transforms/HandshakeOptimizeBitwidths.cpp @@ -1366,19 +1366,88 @@ struct ArithShrUIFW : OpRewritePattern { NameAnalysis &namer; }; +/// Optimizes signed right-shifts with a constant as a forward pass. +struct ArithShrSIFW : OpRewritePattern { + ArithShrSIFW(Pass::Statistic &bitwidthReduced, MLIRContext *ctx, + NameAnalysis &namer) + : OpRewritePattern(ctx), bitwidthReduced(bitwidthReduced), namer(namer) {} + + LogicalResult matchAndRewrite(handshake::ShRSIOp op, + PatternRewriter &rewriter) const override { + auto [lhs, lhsExt] = getMinimalValueWithExtType(op.getLhs()); + unsigned inputBitwidth = lhs.getType().getDataBitWidth(); + unsigned currentBitwidth = op.getType().getDataBitWidth(); + if (inputBitwidth >= currentBitwidth) + return failure(); + + assert(lhsExt != ExtType::NONE && "expected an extension"); + APInt value; + Value constantControl; + { + auto constantOp = op.getRhs().getDefiningOp(); + if (!constantOp) + return failure(); + value = cast(constantOp.getValue()).getValue(); + constantControl = constantOp.getCtrl(); + } + + // Other pattern (such as canonicalization pattern) should fold this case + // to a useful constant instead + if (value.uge(currentBitwidth)) + return failure(); + + if (lhsExt == ExtType::ZEXT) { + // We use a generic canonicalization pattern that should fold this into + // an unsigned shift-right instead. + return failure(); + } + + // SEXT case. + if (value.ult(inputBitwidth)) { + // c is less than the input bitwidth, meaning other bits from the input + // besides the sign-bit are preserved in the output. + + modArithOp(op, {lhs, lhsExt}, {op.getRhs(), ExtType::NONE}, inputBitwidth, + ExtType::SEXT, rewriter, namer); + ++bitwidthReduced; + return success(); + } + + // Our shift amount is larger than the input bitwidth but the input + // bitwidth is sign-extended. The only thing that remains from the input + // is the sign-bit. + Value inputBWM1 = rewriter.create( + op.getLoc(), + rewriter.getIntegerAttr(lhs.getType().getDataType(), inputBitwidth - 1), + constantControl); + // Shift away all values of lhs other than the sign-bit. + ChannelVal signBit = + rewriter.create(op.getLoc(), lhs, inputBWM1); + // Fill remaining sign-bit copies. + rewriter.replaceOpWithNewOp(op, op.getType(), signBit); + ++bitwidthReduced; + return success(); + } + +private: + Pass::Statistic &bitwidthReduced; + /// A reference to the pass's name analysis. + NameAnalysis &namer; +}; + /// Optimizes the bitwidth of shift-type operations. The first template /// parameter is meant to be either handshake::ShLIOp, handshake::ShRSIOp, or /// handshake::ShRUIOp. In both modes (forward and backward), the matched /// operation's bitwidth may only be reduced when the data operand is shifted by /// a known constant amount. template -struct ArithShift : public OpRewritePattern { +struct ArithShiftBW : public OpRewritePattern { using OpRewritePattern::OpRewritePattern; - ArithShift(Pass::Statistic &bitwidthReduced, bool forward, MLIRContext *ctx, - NameAnalysis &namer) + ArithShiftBW(Pass::Statistic &bitwidthReduced, MLIRContext *ctx, + NameAnalysis &namer) : OpRewritePattern(ctx), bitwidthReduced(bitwidthReduced), - namer(namer), forward(forward) {} + namer(namer) {} LogicalResult matchAndRewrite(Op op, PatternRewriter &rewriter) const override { @@ -1396,56 +1465,25 @@ struct ArithShift : public OpRewritePattern { if (Operation *defOp = minShiftBy.getDefiningOp()) if (auto cstOp = dyn_cast(defOp)) { cstVal = (unsigned)cast(cstOp.getValue()).getInt(); - if (forward) { - optWidth = minToShift.getType().getDataBitWidth(); - if (!isRightShift) - optWidth += cstVal; - } else { - optWidth = getUsefulResultWidth(op.getResult()); - if (isRightShift) - optWidth += cstVal; - } + optWidth = getUsefulResultWidth(op.getResult()); + if (isRightShift) + optWidth += cstVal; } if (optWidth >= resWidth) return failure(); - if (forward) { - // Create a new operation as well as appropriate bitwidth modification - // operations to keep the IR valid - Value newToShift = - modBitWidth({minToShift, extToShift}, optWidth, rewriter); - Value newShifyBy = - modBitWidth({minShiftBy, ExtType::ZEXT}, optWidth, rewriter); - rewriter.setInsertionPoint(op); - auto newOp = rewriter.create(op.getLoc(), newToShift.getType(), - newToShift, newShifyBy); - ChannelVal newRes = newOp.getResult(); - if (isRightShift) - // In the case of a right shift, we first truncate the result of the - // newly inserted shift operation to discard high-significance bits that - // we know are 0s, then extend the result back to satisfy the users of - // the original operation's result - newRes = modBitWidth({newRes, extToShift}, optWidth - cstVal, rewriter); - Value modRes = modBitWidth({newRes, extToShift}, resWidth, rewriter); - inheritBB(op, newOp); - - // Replace uses of the original operation's result with the result of the - // optimized operation we just created - rewriter.replaceOp(op, modRes); - } else { - ChannelVal modToShift = minToShift; - if (!isRightShift) { - // In the case of a left shift, we first truncate the shifted integer to - // discard high-significance bits that were discarded in the result, - // then extend back to satisfy the users of the original integer - unsigned requiredToShiftWidth = optWidth - std::min(cstVal, optWidth); - modToShift = modBitWidth({minToShift, extToShift}, requiredToShiftWidth, - rewriter); - } - modArithOp(op, {modToShift, extToShift}, {minShiftBy, ExtType::ZEXT}, - optWidth, extToShift, rewriter, namer); + ChannelVal modToShift = minToShift; + if (!isRightShift) { + // In the case of a left shift, we first truncate the shifted integer to + // discard high-significance bits that were discarded in the result, + // then extend back to satisfy the users of the original integer + unsigned requiredToShiftWidth = optWidth - std::min(cstVal, optWidth); + modToShift = + modBitWidth({minToShift, extToShift}, requiredToShiftWidth, rewriter); } + modArithOp(op, {modToShift, extToShift}, {minShiftBy, ExtType::ZEXT}, + optWidth, extToShift, rewriter, namer); ++bitwidthReduced; return success(); } @@ -1872,19 +1910,20 @@ void HandshakeOptimizeBitwidthsPass::addArithPatterns( // is dangerous if the shift is used as multiplication. // Therefore, removing "ArithShift" from the patterns for // now - patterns.add, ArithSelect>( - bitwidthReduced, forward, ctx, getAnalysis()); - if (!forward) - patterns.add>(bitwidthReduced, forward, ctx, - getAnalysis()); + patterns.add(bitwidthReduced, forward, ctx, + getAnalysis()); + if (forward) + patterns.add(bitwidthReduced, ctx, + getAnalysis()); else - patterns.add(bitwidthReduced, ctx, - getAnalysis()); + patterns.add, + ArithShiftBW>(bitwidthReduced, ctx, + getAnalysis()); patterns.add(bitwidthReduced, ctx, getAnalysis()); - handshake::ExtSIOp::getCanonicalizationPatterns(patterns, ctx); - handshake::ExtUIOp::getCanonicalizationPatterns(patterns, ctx); + ctx->getLoadedDialect() + ->getCanonicalizationPatterns(patterns); } void HandshakeOptimizeBitwidthsPass::addHandshakeDataPatterns( From 35a57015b7a0e46fed21fefd311572d5642ca30d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Markus=20B=C3=B6ck?= Date: Sun, 29 Mar 2026 14:35:00 +0200 Subject: [PATCH 2/3] address review comments from shrui --- lib/Transforms/HandshakeOptimizeBitwidths.cpp | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/lib/Transforms/HandshakeOptimizeBitwidths.cpp b/lib/Transforms/HandshakeOptimizeBitwidths.cpp index 1900c72af2..b8f8571fcd 100644 --- a/lib/Transforms/HandshakeOptimizeBitwidths.cpp +++ b/lib/Transforms/HandshakeOptimizeBitwidths.cpp @@ -1381,19 +1381,19 @@ struct ArithShrSIFW : OpRewritePattern { return failure(); assert(lhsExt != ExtType::NONE && "expected an extension"); - APInt value; + APInt numOfShiftPositions; Value constantControl; { auto constantOp = op.getRhs().getDefiningOp(); if (!constantOp) return failure(); - value = cast(constantOp.getValue()).getValue(); + numOfShiftPositions = cast(constantOp.getValue()).getValue(); constantControl = constantOp.getCtrl(); } // Other pattern (such as canonicalization pattern) should fold this case // to a useful constant instead - if (value.uge(currentBitwidth)) + if (numOfShiftPositions.uge(currentBitwidth)) return failure(); if (lhsExt == ExtType::ZEXT) { @@ -1403,7 +1403,11 @@ struct ArithShrSIFW : OpRewritePattern { } // SEXT case. - if (value.ult(inputBitwidth)) { + // At this point we can reduce the shift to be performed on the lower + // input bitwidth. + // Additional extensions are left to be folded into other operations if + // redundant. + if (numOfShiftPositions.ult(inputBitwidth)) { // c is less than the input bitwidth, meaning other bits from the input // besides the sign-bit are preserved in the output. From c98cd6a9e663ba9ed1fd6c21ae91bcef330861fc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Markus=20B=C3=B6ck?= Date: Sun, 29 Mar 2026 14:36:36 +0200 Subject: [PATCH 3/3] fix test --- .../HandshakeOptimizeBitwidths/arith-forward.mlir | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/test/Transforms/HandshakeOptimizeBitwidths/arith-forward.mlir b/test/Transforms/HandshakeOptimizeBitwidths/arith-forward.mlir index 94d72cf12a..61b2e167b3 100644 --- a/test/Transforms/HandshakeOptimizeBitwidths/arith-forward.mlir +++ b/test/Transforms/HandshakeOptimizeBitwidths/arith-forward.mlir @@ -152,10 +152,10 @@ handshake.func @shliFW(%arg0: !handshake.channel, %start: !handshake.contro // CHECK-LABEL: handshake.func @shrsiFW( // CHECK-SAME: %[[VAL_0:.*]]: !handshake.channel, // CHECK-SAME: %[[VAL_1:.*]]: !handshake.control<>, ...) -> !handshake.channel attributes {argNames = ["arg0", "start"], resNames = ["out0"]} { -// CHECK: %[[VAL_2:.*]] = constant %[[VAL_1]] {value = 4 : i16} : <>, -// CHECK: %[[VAL_3:.*]] = shrsi %[[VAL_0]], %[[VAL_2]] : -// CHECK: %[[VAL_4:.*]] = trunci %[[VAL_3]] : to -// CHECK: %[[VAL_5:.*]] = extsi %[[VAL_4]] : to +// CHECK: %[[VAL_2:.*]] = constant %[[VAL_1]] {value = 4 : i32} : <>, +// CHECK: %[[VAL_3:.*]] = trunci %[[VAL_2]] : to +// CHECK: %[[VAL_4:.*]] = shrsi %[[VAL_0]], %[[VAL_3]] : +// CHECK: %[[VAL_5:.*]] = extsi %[[VAL_4]] : to // CHECK: end %[[VAL_5]] : // CHECK: } handshake.func @shrsiFW(%arg0: !handshake.channel, %start: !handshake.control<>) -> !handshake.channel {