From c4467d978f28009d9cb9ecf2a07ddb5da5dfe2be Mon Sep 17 00:00:00 2001 From: Angelo Matni Date: Fri, 17 Jul 2026 12:13:37 -0700 Subject: [PATCH] [opt] Fix CanonicalizePass clamping: create Literal from known bits if dne If the clamped value was not a Literal but was known, MaybeCanonicalizeClamp previously assumed it was a Literal and crashed. Now we construct a Literal from the Bits returned by the query engine. PiperOrigin-RevId: 949696244 --- xls/passes/BUILD | 1 + xls/passes/canonicalization_pass.cc | 87 ++++++++++++++---------- xls/passes/canonicalization_pass_test.cc | 52 ++++++++++++++ 3 files changed, 105 insertions(+), 35 deletions(-) diff --git a/xls/passes/BUILD b/xls/passes/BUILD index 75b509f0e3..e874c3f348 100644 --- a/xls/passes/BUILD +++ b/xls/passes/BUILD @@ -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", diff --git a/xls/passes/canonicalization_pass.cc b/xls/passes/canonicalization_pass.cc index 6ef2ed80e5..43a4143caa 100644 --- a/xls/passes/canonicalization_pass.cc +++ b/xls/passes/canonicalization_pass.cc @@ -16,6 +16,7 @@ #include #include +#include #include #include "absl/log/log.h" @@ -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> AreSequentialConstants( + Node* m, Node* n, QueryEngine& query_engine) { std::optional m_bits = query_engine.KnownValueAsBits(m); std::optional 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 AreEqualConstants(Node* m, Node* n, + QueryEngine& query_engine) { std::optional m_bits = query_engine.KnownValueAsBits(m); std::optional 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 AsLiteralOrMakeLiteral(Node* candidate, + const Bits& bits) { + if (candidate->Is() && candidate->As()->value().IsBits() && + candidate->As()->value().bits() == bits) { + return candidate->As(); } - return bits_ops::UEqual(*m_bits, *n_bits); + return candidate->function_base()->MakeNode(candidate->loc(), + Value(bits)); } // Change clamps to high or low values to a canonical form: @@ -104,42 +116,47 @@ absl::StatusOr 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 b_d_equal = AreEqualConstants(b, d, query_engine); + std::optional> b_c_sequential = + AreSequentialConstants(b, c, query_engine); + std::optional> b_d_sequential = + AreSequentialConstants(b, d, query_engine); + std::optional> c_b_sequential = + AreSequentialConstants(c, b, query_engine); + std::optional> 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(); - } 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(); - } 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(); - } 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(); - } 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(); - } 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(); + XLS_ASSIGN_OR_RETURN(k, AsLiteralOrMakeLiteral(d, b_d_sequential->second)); } + if (is_clamp_high || is_clamp_low) { // Create an expression: // diff --git a/xls/passes/canonicalization_pass_test.cc b/xls/passes/canonicalization_pass_test.cc index 27fb038447..4847b39006 100644 --- a/xls/passes/canonicalization_pass_test.cc +++ b/xls/passes/canonicalization_pass_test.cc @@ -42,6 +42,7 @@ #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; @@ -49,6 +50,7 @@ namespace xls { namespace { using ::absl_testing::IsOkAndHolds; +using solvers::ScopedVerifyEquivalence; using ::testing::ElementsAre; using ::testing::IsEmpty; using ::testing::Optional; @@ -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"(