diff --git a/xls/passes/BUILD b/xls/passes/BUILD index 0730159ffc..53779099fe 100644 --- a/xls/passes/BUILD +++ b/xls/passes/BUILD @@ -904,6 +904,7 @@ xls_pass( "@abseil-cpp//absl/log:check", "@abseil-cpp//absl/status", "@abseil-cpp//absl/status:statusor", + "@abseil-cpp//absl/strings", "@abseil-cpp//absl/strings:str_format", "@abseil-cpp//absl/types:span", ], diff --git a/xls/passes/strength_reduction_pass.cc b/xls/passes/strength_reduction_pass.cc index f5f26ca98b..94316d7500 100644 --- a/xls/passes/strength_reduction_pass.cc +++ b/xls/passes/strength_reduction_pass.cc @@ -27,6 +27,7 @@ #include "absl/status/status.h" #include "absl/status/statusor.h" #include "absl/strings/str_format.h" +#include "absl/strings/str_join.h" #include "absl/types/span.h" #include "xls/common/status/ret_check.h" #include "xls/common/status/status_macros.h" @@ -153,6 +154,37 @@ absl::StatusOr MaybeSinkOperationIntoSelect( return false; } +absl::StatusOr SplitBinOp(Node* node, Op op, + absl::Span split_points) { + Node* lhs = node->operand(0); + Node* rhs = node->operand(1); + int64_t width = node->BitCountOrDie(); + + std::vector bounds; + bounds.push_back(0); + bounds.insert(bounds.end(), split_points.begin(), split_points.end()); + bounds.push_back(width); + + std::vector parts; + for (int64_t i = bounds.size() - 2; i >= 0; --i) { + int64_t seg_start = bounds[i]; + int64_t seg_width = bounds[i + 1] - bounds[i]; + + XLS_ASSIGN_OR_RETURN(Node * lhs_slice, + node->function_base()->MakeNode( + node->loc(), lhs, seg_start, seg_width)); + XLS_ASSIGN_OR_RETURN(Node * rhs_slice, + node->function_base()->MakeNode( + node->loc(), rhs, seg_start, seg_width)); + XLS_ASSIGN_OR_RETURN(Node * part, + node->function_base()->MakeNode( + node->loc(), lhs_slice, rhs_slice, op)); + parts.push_back(part); + } + + return node->function_base()->MakeNode(node->loc(), parts); +} + // Attempts to strength-reduce the given node. Returns true if successful. absl::StatusOr StrengthReduceNode(Node* node, const QueryEngine& query_engine, @@ -611,92 +643,51 @@ absl::StatusOr StrengthReduceNode(Node* node, return ternary_tree->Get({}); }; if (SplitsEnabled(opt_level) && node->op() == Op::kAdd) { - Node* add = node; Node* lhs = node->operand(0); Node* rhs = node->operand(1); TernaryVector propagate_carry = ternary_ops::Or(get_ternary(lhs), get_ternary(rhs)); - if (auto non_propagate_it = - absl::c_find(propagate_carry, TernaryValue::kKnownZero); - non_propagate_it != propagate_carry.end() && - non_propagate_it != propagate_carry.end() - 1) { - int64_t split_point = - std::distance(propagate_carry.begin(), non_propagate_it) + 1; - VLOG(2) << "Add cannot propagate carries to bit " << split_point - << "; replacing with two adders: " << add->ToString(); - - XLS_ASSIGN_OR_RETURN(Node * trailing_lhs, - node->function_base()->MakeNode( - node->loc(), lhs, 0, split_point)); - XLS_ASSIGN_OR_RETURN(Node * leading_lhs, - node->function_base()->MakeNode( - node->loc(), lhs, split_point, - lhs->BitCountOrDie() - split_point)); - XLS_ASSIGN_OR_RETURN(Node * trailing_rhs, - node->function_base()->MakeNode( - node->loc(), rhs, 0, split_point)); - XLS_ASSIGN_OR_RETURN(Node * leading_rhs, - node->function_base()->MakeNode( - node->loc(), rhs, split_point, - rhs->BitCountOrDie() - split_point)); - XLS_ASSIGN_OR_RETURN(Node * leading_add, - add->function_base()->MakeNode( - add->loc(), leading_lhs, leading_rhs, Op::kAdd)); - XLS_ASSIGN_OR_RETURN( - Node * trailing_add, - add->function_base()->MakeNode(add->loc(), trailing_lhs, - trailing_rhs, Op::kAdd)); - XLS_RETURN_IF_ERROR( - add->ReplaceUsesWithNew( - absl::Span{leading_add, trailing_add}) - .status()); + std::vector split_points; + for (int64_t i = 0; i < propagate_carry.size() - 1; ++i) { + if (propagate_carry[i] == TernaryValue::kKnownZero) { + split_points.push_back(i + 1); + } + } + + if (!split_points.empty()) { + VLOG(2) << "Add cannot propagate carries at bits " + << absl::StrJoin(split_points, ", ") << "; replacing with " + << split_points.size() + 1 << " adders: " << node->ToString(); + XLS_ASSIGN_OR_RETURN(Node * split_node, + SplitBinOp(node, Op::kAdd, split_points)); + XLS_RETURN_IF_ERROR(node->ReplaceUsesWith(split_node)); return true; } } if (SplitsEnabled(opt_level) && node->op() == Op::kSub) { - Node* sub = node; Node* lhs = node->operand(0); Node* rhs = node->operand(1); TernaryVector borrow_stop = ternary_ops::And(get_ternary(lhs), ternary_ops::Not(get_ternary(rhs))); - if (auto borrow_stop_it = - absl::c_find(borrow_stop, TernaryValue::kKnownOne); - borrow_stop_it != borrow_stop.end() && - borrow_stop_it != borrow_stop.end() - 1) { - int64_t split_point = - std::distance(borrow_stop.begin(), borrow_stop_it) + 1; - VLOG(2) << "Sub cannot propagate borrows to bit " << split_point - << "; replacing with two subtractors: " << sub->ToString(); - - XLS_ASSIGN_OR_RETURN(Node * trailing_lhs, - node->function_base()->MakeNode( - node->loc(), lhs, 0, split_point)); - XLS_ASSIGN_OR_RETURN(Node * leading_lhs, - node->function_base()->MakeNode( - node->loc(), lhs, split_point, - lhs->BitCountOrDie() - split_point)); - XLS_ASSIGN_OR_RETURN(Node * trailing_rhs, - node->function_base()->MakeNode( - node->loc(), rhs, 0, split_point)); - XLS_ASSIGN_OR_RETURN(Node * leading_rhs, - node->function_base()->MakeNode( - node->loc(), rhs, split_point, - rhs->BitCountOrDie() - split_point)); - XLS_ASSIGN_OR_RETURN(Node * leading_sub, - sub->function_base()->MakeNode( - sub->loc(), leading_lhs, leading_rhs, Op::kSub)); - XLS_ASSIGN_OR_RETURN( - Node * trailing_sub, - sub->function_base()->MakeNode(sub->loc(), trailing_lhs, - trailing_rhs, Op::kSub)); - XLS_RETURN_IF_ERROR( - sub->ReplaceUsesWithNew( - absl::Span{leading_sub, trailing_sub}) - .status()); + std::vector split_points; + for (int64_t i = 0; i < borrow_stop.size() - 1; ++i) { + if (borrow_stop[i] == TernaryValue::kKnownOne) { + split_points.push_back(i + 1); + } + } + + if (!split_points.empty()) { + VLOG(2) << "Sub cannot propagate borrows at bits " + << absl::StrJoin(split_points, ", ") << "; replacing with " + << split_points.size() + 1 + << " subtractors: " << node->ToString(); + XLS_ASSIGN_OR_RETURN(Node * split_node, + SplitBinOp(node, Op::kSub, split_points)); + XLS_RETURN_IF_ERROR(node->ReplaceUsesWith(split_node)); return true; } } diff --git a/xls/passes/strength_reduction_pass_test.cc b/xls/passes/strength_reduction_pass_test.cc index 2ef0bc5a8a..895fe751b8 100644 --- a/xls/passes/strength_reduction_pass_test.cc +++ b/xls/passes/strength_reduction_pass_test.cc @@ -914,8 +914,10 @@ TEST_F(StrengthReductionPassTest, AdderSplitOnNoCarryPropagateMultiple) { ASSERT_THAT(Run(f), IsOkAndHolds(true)); EXPECT_THAT( f->return_value(), - m::Concat(m::Add(m::BitSlice(lhs.node(), /*start=*/2, /*width=*/3), - m::BitSlice(rhs.node(), /*start=*/2, /*width=*/3)), + m::Concat(m::Add(m::BitSlice(lhs.node(), /*start=*/4, /*width=*/1), + m::BitSlice(rhs.node(), /*start=*/4, /*width=*/1)), + m::Add(m::BitSlice(lhs.node(), /*start=*/2, /*width=*/2), + m::BitSlice(rhs.node(), /*start=*/2, /*width=*/2)), m::Add(m::BitSlice(lhs.node(), /*start=*/0, /*width=*/2), m::BitSlice(rhs.node(), /*start=*/0, /*width=*/2)))); } @@ -953,7 +955,30 @@ TEST_F(StrengthReductionPassTest, SubSplitOnNoBorrowPropagateBit2) { m::BitSlice(rhs.node(), /*start=*/0, /*width=*/3)))); } -TEST_F(StrengthReductionPassTest, SubSplitOnNoBorrowPropagateLowestBit) { +TEST_F(StrengthReductionPassTest, AdderSplitOnNoCarryPropagateAdjacent) { + auto p = CreatePackage(); + FunctionBuilder fb(TestName(), p.get()); + BValue p0 = fb.Param("p0", p->GetBitsType(5)); + BValue p1 = fb.Param("p1", p->GetBitsType(5)); + // Guarantees bit 1 and bit 2 of both operands are 0 -> + // propagate_carry[1]=0, propagate_carry[2]=0 + BValue lhs = fb.And(p0, fb.Literal(UBits(0b11001, 5))); + BValue rhs = fb.And(p1, fb.Literal(UBits(0b11001, 5))); + fb.Add(lhs, rhs); + XLS_ASSERT_OK_AND_ASSIGN(Function * f, fb.Build()); + ScopedVerifyEquivalence sve(f); + ASSERT_THAT(Run(f), IsOkAndHolds(true)); + EXPECT_THAT( + f->return_value(), + m::Concat(m::Add(m::BitSlice(lhs.node(), /*start=*/3, /*width=*/2), + m::BitSlice(rhs.node(), /*start=*/3, /*width=*/2)), + m::Add(m::BitSlice(lhs.node(), /*start=*/2, /*width=*/1), + m::BitSlice(rhs.node(), /*start=*/2, /*width=*/1)), + m::Add(m::BitSlice(lhs.node(), /*start=*/0, /*width=*/2), + m::BitSlice(rhs.node(), /*start=*/0, /*width=*/2)))); +} + +TEST_F(StrengthReductionPassTest, SubSplitOnNoBorrowPropagateMultiple) { auto p = CreatePackage(); FunctionBuilder fb(TestName(), p.get()); BValue p0 = fb.Param("p0", p->GetBitsType(5)); @@ -965,12 +990,15 @@ TEST_F(StrengthReductionPassTest, SubSplitOnNoBorrowPropagateLowestBit) { fb.Subtract(lhs, rhs); XLS_ASSERT_OK_AND_ASSIGN(Function * f, fb.Build()); ScopedVerifyEquivalence sve(f); - // Bit 1 is the lowest bit where lhs=1, rhs=0. Should split at bit 1+1=2. + // Bits 1 and 3 are bits where lhs=1, rhs=0. + // Should split at bits 1+1=2 and 3+1=4. ASSERT_THAT(Run(f), IsOkAndHolds(true)); EXPECT_THAT( f->return_value(), - m::Concat(m::Sub(m::BitSlice(lhs.node(), /*start=*/2, /*width=*/3), - m::BitSlice(rhs.node(), /*start=*/2, /*width=*/3)), + m::Concat(m::Sub(m::BitSlice(lhs.node(), /*start=*/4, /*width=*/1), + m::BitSlice(rhs.node(), /*start=*/4, /*width=*/1)), + m::Sub(m::BitSlice(lhs.node(), /*start=*/2, /*width=*/2), + m::BitSlice(rhs.node(), /*start=*/2, /*width=*/2)), m::Sub(m::BitSlice(lhs.node(), /*start=*/0, /*width=*/2), m::BitSlice(rhs.node(), /*start=*/0, /*width=*/2)))); }