Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions xls/passes/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -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",
],
Expand Down
131 changes: 61 additions & 70 deletions xls/passes/strength_reduction_pass.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -153,6 +154,37 @@ absl::StatusOr<bool> MaybeSinkOperationIntoSelect(
return false;
}

absl::StatusOr<Node*> SplitBinOp(Node* node, Op op,
absl::Span<const int64_t> split_points) {
Node* lhs = node->operand(0);
Node* rhs = node->operand(1);
int64_t width = node->BitCountOrDie();

std::vector<int64_t> bounds;
bounds.push_back(0);
bounds.insert(bounds.end(), split_points.begin(), split_points.end());
bounds.push_back(width);

std::vector<Node*> 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<BitSlice>(
node->loc(), lhs, seg_start, seg_width));
XLS_ASSIGN_OR_RETURN(Node * rhs_slice,
node->function_base()->MakeNode<BitSlice>(
node->loc(), rhs, seg_start, seg_width));
XLS_ASSIGN_OR_RETURN(Node * part,
node->function_base()->MakeNode<BinOp>(
node->loc(), lhs_slice, rhs_slice, op));
parts.push_back(part);
}

return node->function_base()->MakeNode<Concat>(node->loc(), parts);
}

// Attempts to strength-reduce the given node. Returns true if successful.
absl::StatusOr<bool> StrengthReduceNode(Node* node,
const QueryEngine& query_engine,
Expand Down Expand Up @@ -611,92 +643,51 @@ absl::StatusOr<bool> 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<BitSlice>(
node->loc(), lhs, 0, split_point));
XLS_ASSIGN_OR_RETURN(Node * leading_lhs,
node->function_base()->MakeNode<BitSlice>(
node->loc(), lhs, split_point,
lhs->BitCountOrDie() - split_point));
XLS_ASSIGN_OR_RETURN(Node * trailing_rhs,
node->function_base()->MakeNode<BitSlice>(
node->loc(), rhs, 0, split_point));
XLS_ASSIGN_OR_RETURN(Node * leading_rhs,
node->function_base()->MakeNode<BitSlice>(
node->loc(), rhs, split_point,
rhs->BitCountOrDie() - split_point));

XLS_ASSIGN_OR_RETURN(Node * leading_add,
add->function_base()->MakeNode<BinOp>(
add->loc(), leading_lhs, leading_rhs, Op::kAdd));
XLS_ASSIGN_OR_RETURN(
Node * trailing_add,
add->function_base()->MakeNode<BinOp>(add->loc(), trailing_lhs,
trailing_rhs, Op::kAdd));
XLS_RETURN_IF_ERROR(
add->ReplaceUsesWithNew<Concat>(
absl::Span<Node* const>{leading_add, trailing_add})
.status());
std::vector<int64_t> 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<BitSlice>(
node->loc(), lhs, 0, split_point));
XLS_ASSIGN_OR_RETURN(Node * leading_lhs,
node->function_base()->MakeNode<BitSlice>(
node->loc(), lhs, split_point,
lhs->BitCountOrDie() - split_point));
XLS_ASSIGN_OR_RETURN(Node * trailing_rhs,
node->function_base()->MakeNode<BitSlice>(
node->loc(), rhs, 0, split_point));
XLS_ASSIGN_OR_RETURN(Node * leading_rhs,
node->function_base()->MakeNode<BitSlice>(
node->loc(), rhs, split_point,
rhs->BitCountOrDie() - split_point));

XLS_ASSIGN_OR_RETURN(Node * leading_sub,
sub->function_base()->MakeNode<BinOp>(
sub->loc(), leading_lhs, leading_rhs, Op::kSub));
XLS_ASSIGN_OR_RETURN(
Node * trailing_sub,
sub->function_base()->MakeNode<BinOp>(sub->loc(), trailing_lhs,
trailing_rhs, Op::kSub));
XLS_RETURN_IF_ERROR(
sub->ReplaceUsesWithNew<Concat>(
absl::Span<Node* const>{leading_sub, trailing_sub})
.status());
std::vector<int64_t> 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;
}
}
Expand Down
40 changes: 34 additions & 6 deletions xls/passes/strength_reduction_pass_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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))));
}
Expand Down Expand Up @@ -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));
Expand All @@ -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))));
}
Expand Down
Loading