Skip to content
Merged
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 @@ -403,6 +403,7 @@ cc_test(
"//xls/ir:ir_test_base",
"//xls/ir:op",
"//xls/ir:value",
"//xls/solvers:ir_equivalence_testutils",
"@abseil-cpp//absl/log",
"@abseil-cpp//absl/status:statusor",
"@abseil-cpp//absl/strings",
Expand Down
87 changes: 52 additions & 35 deletions xls/passes/canonicalization_pass.cc
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@

#include <cstdint>
#include <optional>
#include <utility>
#include <vector>

#include "absl/log/log.h"
Expand All @@ -40,34 +41,45 @@ namespace xls {

namespace {

// Returns true if 'm' and 'n' are both constants whose bits values are
// sequential unsigned values (m + 1 = n).
bool AreSequentialConstants(Node* m, Node* n, QueryEngine& query_engine) {
if (!m->GetType()->IsBits() || !n->GetType()->IsBits()) {
return false;
}
// Returns the Bits values of 'm' and 'n' if they are both constants whose bits
// values are sequential unsigned values (m + 1 = n).
std::optional<std::pair<Bits, Bits>> AreSequentialConstants(
Node* m, Node* n, QueryEngine& query_engine) {
std::optional<Bits> m_bits = query_engine.KnownValueAsBits(m);
std::optional<Bits> n_bits = query_engine.KnownValueAsBits(n);
if (!m_bits.has_value() || !n_bits.has_value()) {
return false;
return std::nullopt;
}
// Zero extend before adding one to avoid overflow.
return bits_ops::UEqual(bits_ops::Increment(bits_ops::ZeroExtend(
*m_bits, m_bits->bit_count() + 1)),
*n_bits);
if (bits_ops::UEqual(bits_ops::Increment(bits_ops::ZeroExtend(
*m_bits, m_bits->bit_count() + 1)),
*n_bits)) {
return std::make_pair(std::move(*m_bits), std::move(*n_bits));
}
return std::nullopt;
}

// Returns true if 'm' and 'n' are both constants whose bits values are equal.
bool AreEqualConstants(Node* m, Node* n, QueryEngine& query_engine) {
if (!m->GetType()->IsBits() || !n->GetType()->IsBits()) {
return false;
}
// Returns the Bits value of 'm' and 'n' if they are both constants whose bits
// values are equal.
std::optional<Bits> AreEqualConstants(Node* m, Node* n,
QueryEngine& query_engine) {
std::optional<Bits> m_bits = query_engine.KnownValueAsBits(m);
std::optional<Bits> n_bits = query_engine.KnownValueAsBits(n);
if (!m_bits.has_value() || !n_bits.has_value()) {
return false;
if (!m_bits.has_value() || !n_bits.has_value() ||
!bits_ops::UEqual(*m_bits, *n_bits)) {
return std::nullopt;
}
return *m_bits;
}

absl::StatusOr<Literal*> AsLiteralOrMakeLiteral(Node* candidate,
const Bits& bits) {
if (candidate->Is<Literal>() && candidate->As<Literal>()->value().IsBits() &&
candidate->As<Literal>()->value().bits() == bits) {
return candidate->As<Literal>();
}
return bits_ops::UEqual(*m_bits, *n_bits);
return candidate->function_base()->MakeNode<Literal>(candidate->loc(),
Value(bits));
}

// Change clamps to high or low values to a canonical form:
Expand Down Expand Up @@ -104,42 +116,47 @@ absl::StatusOr<bool> MaybeCanonicalizeClamp(Node* n,
Literal* k = nullptr;
bool is_clamp_low = false;
bool is_clamp_high = false;
if (cmp == Op::kUGt && a == d && AreSequentialConstants(b, c, query_engine)) {
std::optional<Bits> b_d_equal = AreEqualConstants(b, d, query_engine);
std::optional<std::pair<Bits, Bits>> b_c_sequential =
AreSequentialConstants(b, c, query_engine);
std::optional<std::pair<Bits, Bits>> b_d_sequential =
AreSequentialConstants(b, d, query_engine);
std::optional<std::pair<Bits, Bits>> c_b_sequential =
AreSequentialConstants(c, b, query_engine);
std::optional<std::pair<Bits, Bits>> d_b_sequential =
AreSequentialConstants(d, b, query_engine);
if (cmp == Op::kUGt && a == d && b_c_sequential.has_value()) {
// a cmp b ? c : d
// (i) x > K - 1 ? K : x => x > K ? K : x
is_clamp_high = true;
k = c->As<Literal>();
} else if (cmp == Op::kULt && a == c &&
AreEqualConstants(b, d, query_engine)) {
XLS_ASSIGN_OR_RETURN(k, AsLiteralOrMakeLiteral(c, b_c_sequential->second));
} else if (cmp == Op::kULt && a == c && b_d_equal.has_value()) {
// a cmp b ? c : d
// (ii) x < K ? x : K => x > K ? K : x
is_clamp_high = true;
k = d->As<Literal>();
} else if (cmp == Op::kULt && a == c &&
AreSequentialConstants(d, b, query_engine)) {
XLS_ASSIGN_OR_RETURN(k, AsLiteralOrMakeLiteral(d, *b_d_equal));
} else if (cmp == Op::kULt && a == c && d_b_sequential.has_value()) {
// a cmp b ? c : d
// (iii) x < K + 1 ? x : K => x > K ? K : x
is_clamp_high = true;
k = d->As<Literal>();
} else if (cmp == Op::kULt && a == d &&
AreSequentialConstants(c, b, query_engine)) {
XLS_ASSIGN_OR_RETURN(k, AsLiteralOrMakeLiteral(d, d_b_sequential->first));
} else if (cmp == Op::kULt && a == d && c_b_sequential.has_value()) {
// a cmp b ? c : d
// (iv) x < K + 1 ? K : x => x < K ? K : x
is_clamp_low = true;
k = c->As<Literal>();
} else if (cmp == Op::kUGt && a == c &&
AreEqualConstants(b, d, query_engine)) {
XLS_ASSIGN_OR_RETURN(k, AsLiteralOrMakeLiteral(c, c_b_sequential->first));
} else if (cmp == Op::kUGt && a == c && b_d_equal.has_value()) {
// a cmp b ? c : d
// (v) x > K ? x : K => x < K ? K : x
is_clamp_low = true;
k = d->As<Literal>();
} else if (cmp == Op::kUGt && a == c &&
AreSequentialConstants(b, d, query_engine)) {
XLS_ASSIGN_OR_RETURN(k, AsLiteralOrMakeLiteral(d, *b_d_equal));
} else if (cmp == Op::kUGt && a == c && b_d_sequential.has_value()) {
// a cmp b ? c : d
// (vi) x > K - 1 ? x : K => x < K ? K : x
is_clamp_low = true;
k = d->As<Literal>();
XLS_ASSIGN_OR_RETURN(k, AsLiteralOrMakeLiteral(d, b_d_sequential->second));
}

if (is_clamp_high || is_clamp_low) {
// Create an expression:
//
Expand Down
52 changes: 52 additions & 0 deletions xls/passes/canonicalization_pass_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -42,13 +42,15 @@
#include "xls/ir/value.h"
#include "xls/passes/optimization_pass.h"
#include "xls/passes/pass_base.h"
#include "xls/solvers/ir_equivalence_testutils.h"

namespace m = ::xls::op_matchers;

namespace xls {
namespace {

using ::absl_testing::IsOkAndHolds;
using solvers::ScopedVerifyEquivalence;
using ::testing::ElementsAre;
using ::testing::IsEmpty;
using ::testing::Optional;
Expand Down Expand Up @@ -259,6 +261,56 @@ TEST_F(CanonicalizePassTest, ExhaustiveClampTest) {
}
}

TEST_F(CanonicalizePassTest, ClampWithKnownValuesNotLiterals) {
auto p = CreatePackage();
FunctionBuilder fb(TestName(), p.get());
BValue x = fb.Param("x", p->GetBitsType(4));
BValue lit_2 = fb.Literal(UBits(2, 3));
BValue lit_3 = fb.Literal(UBits(3, 3));
// Canonicalize x > 2 ? 3 : x => x > 3 ? 3 : x
BValue sel1 = fb.Select(fb.UGt(x, fb.ZeroExtend(lit_2, 4)),
/*cases=*/{x, fb.ZeroExtend(lit_3, 4)});
// Canonicalize x < 3 ? 2 : x => x < 2 ? 2 : x
BValue sel2 = fb.Select(fb.ULt(x, fb.ZeroExtend(lit_3, 4)),
/*cases=*/{x, fb.ZeroExtend(lit_2, 4)});
// Canonicalize x < 2 ? x : 2 => x > 2 ? 2 : x
BValue sel3 = fb.Select(fb.ULt(x, fb.ZeroExtend(lit_2, 4)),
/*cases=*/{fb.ZeroExtend(lit_2, 4), x});
// Canonicalize x < 3 ? x : 2 => x > 2 ? 2 : x
BValue sel4 = fb.Select(fb.ULt(x, fb.ZeroExtend(lit_3, 4)),
/*cases=*/{fb.ZeroExtend(lit_2, 4), x});
// Canonicalize x > 2 ? x : 2 => x < 2 ? 2 : x
BValue sel5 = fb.Select(fb.UGt(x, fb.ZeroExtend(lit_2, 4)),
/*cases=*/{fb.ZeroExtend(lit_2, 4), x});
XLS_ASSERT_OK_AND_ASSIGN(
Function * f,
fb.BuildWithReturnValue(fb.Tuple({sel1, sel2, sel3, sel4, sel5})));
ScopedVerifyEquivalence sve(f);
EXPECT_THAT(Run(p.get()), IsOkAndHolds(true));
EXPECT_THAT(f->return_value(),
m::Tuple(m::Select(m::UGt(m::Param("x"), m::Literal(3)),
/*cases=*/{m::Param("x"), m::Literal(3)}),
m::Select(m::ULt(m::Param("x"), m::Literal(2)),
/*cases=*/{m::Param("x"), m::Literal(2)}),
m::Select(m::UGt(m::Param("x"), m::Literal(2)),
/*cases=*/{m::Param("x"), m::Literal(2)}),
m::Select(m::UGt(m::Param("x"), m::Literal(2)),
/*cases=*/{m::Param("x"), m::Literal(2)}),
m::Select(m::ULt(m::Param("x"), m::Literal(2)),
/*cases=*/{m::Param("x"), m::Literal(2)})));
}

TEST_F(CanonicalizePassTest, DoNotClampWithNonBitValues) {
auto p = CreatePackage();
FunctionBuilder fb(TestName(), p.get());
BValue x = fb.Param("x", p->GetBitsType(4));
BValue select =
fb.Select(fb.UGt(x, fb.Literal(UBits(1, 4))),
{fb.Tuple({x}), fb.Tuple({fb.Literal(UBits(2, 4))})});
XLS_ASSERT_OK(fb.BuildWithReturnValue(select));
EXPECT_THAT(Run(p.get()), IsOkAndHolds(false));
}

TEST_F(CanonicalizePassTest, SelectWithTrivialDefault) {
auto p = CreatePackage();
XLS_ASSERT_OK_AND_ASSIGN(Function * f, ParseFunction(R"(
Expand Down
Loading