diff --git a/xls/passes/arith_simplification_pass.cc b/xls/passes/arith_simplification_pass.cc index b6be6d252f..9773e2b5f9 100644 --- a/xls/passes/arith_simplification_pass.cc +++ b/xls/passes/arith_simplification_pass.cc @@ -1188,6 +1188,30 @@ absl::StatusOr MatchDivByVariablePowerOfTwo( return true; } +struct LeadingZerosTrailingOnes { + int64_t leading_zeros; + int64_t trailing_ones; +}; + +// If the given `bits` is a constant mask of the form 0...01...1, returns the +// number of leading zeros and trailing ones. Otherwise, returns std::nullopt. +std::optional LsbConstantMask(const Bits& bits) { + int64_t trailing_zeros = -1; + int64_t leading_zeros = -1; + int64_t trailing_ones = -1; + bool is_mask = bits.HasSingleRunOfSetBits(&leading_zeros, &trailing_ones, + &trailing_zeros); + + if (!is_mask || trailing_zeros != 0) { + return std::nullopt; + } + + return LeadingZerosTrailingOnes{ + .leading_zeros = leading_zeros, + .trailing_ones = trailing_ones, + }; +} + // MatchArithPatterns matches simple tree patterns to find opportunities // for simplification, such as adding a zero, multiplying by 1, etc. // @@ -1964,71 +1988,80 @@ absl::StatusOr MatchArithPatterns(int64_t opt_level, Node* n, // An unsigned comparison against a constant mask of LSBs (e.g., 0b0001111) // can be simplified: // - // x < 0b0001111 => or_reduce(msb_slice(x)) NOR - // and_reduce(lsb_slice(x)) x > 0b0001111 => or_reduce(msb_slice(x)) + // x < 0b0001111 => or_reduce(msb_slice(x)) NOR and_reduce(lsb_slice(x)) + // x > 0b0001111 => or_reduce(msb_slice(x)) // // x <= 0b0001111 => nor_reduce(msb_slice(x)) // x >= 0b0001111 => or_reduce(msb_slice(x)) OR and_reduce(lsb_slice(x)) - int64_t leading_zeros, trailing_ones; - auto is_constant_mask = [&](Node* node, int64_t* leading_zero_count, - int64_t* trailing_one_count) { + // + // Similarly: + // (x > mask - 1) <=> (x >= mask) + auto get_mask_cmp = + [&](Node* node, + Op op) -> std::optional> { const std::optional constant = query_engine.KnownValueAsBits(node); if (!constant.has_value()) { - return false; + return std::nullopt; } - int64_t trailing_zeros = -1; - if (!constant->HasSingleRunOfSetBits(leading_zero_count, trailing_one_count, - &trailing_zeros)) { - return false; + if (auto mask = LsbConstantMask(*constant); mask.has_value()) { + return std::make_tuple(op, mask->leading_zeros, mask->trailing_ones); + } + if (auto mask = LsbConstantMask(bits_ops::Increment(*constant)); + mask.has_value() && op == Op::kUGt && !constant->IsAllOnes()) { + return std::make_tuple(Op::kUGe, mask->leading_zeros, + mask->trailing_ones); } - return trailing_zeros == 0; + return std::nullopt; }; - if (NarrowingEnabled(opt_level) && IsBitsCompare(n) && - is_constant_mask(n->operand(1), &leading_zeros, &trailing_ones)) { - VLOG(2) << "Found comparison to literal mask; leading zeros: " - << leading_zeros << " trailing ones: " << trailing_ones - << " :: " << n; - switch (n->op()) { - case Op::kULt: { - XLS_ASSIGN_OR_RETURN( - Node * or_red, - OrReduceLeading(n->operand(0), leading_zeros, n->loc())); - XLS_ASSIGN_OR_RETURN( - Node * and_trail, - AndReduceTrailing(n->operand(0), trailing_ones, n->loc())); - std::vector args = {or_red, and_trail}; - XLS_RETURN_IF_ERROR( - n->ReplaceUsesWithNew(args, Op::kNor).status()); - return true; - } - case Op::kUGt: { - XLS_ASSIGN_OR_RETURN( - Node * or_red, - OrReduceLeading(n->operand(0), leading_zeros, n->loc())); - XLS_RETURN_IF_ERROR(n->ReplaceUsesWith(or_red)); - return true; - } - case Op::kULe: { - XLS_ASSIGN_OR_RETURN( - Node * nor_red, - NorReduceLeading(n->operand(0), leading_zeros, n->loc())); - XLS_RETURN_IF_ERROR(n->ReplaceUsesWith(nor_red)); - return true; - } - case Op::kUGe: { - XLS_ASSIGN_OR_RETURN( - Node * or_red, - OrReduceLeading(n->operand(0), leading_zeros, n->loc())); - XLS_ASSIGN_OR_RETURN( - Node * and_trail, - AndReduceTrailing(n->operand(0), trailing_ones, n->loc())); - std::vector args = {or_red, and_trail}; - XLS_RETURN_IF_ERROR( - n->ReplaceUsesWithNew(args, Op::kOr).status()); - return true; + + if (NarrowingEnabled(opt_level) && IsBitsCompare(n)) { + if (auto mask_cmp = get_mask_cmp(n->operand(1), n->op())) { + auto [effective_op, leading_zeros, trailing_ones] = *mask_cmp; + VLOG(2) << "Found comparison to literal mask; leading zeros: " + << leading_zeros << " trailing ones: " << trailing_ones + << " :: " << n; + switch (effective_op) { + case Op::kULt: { + XLS_ASSIGN_OR_RETURN( + Node * or_red, + OrReduceLeading(n->operand(0), leading_zeros, n->loc())); + XLS_ASSIGN_OR_RETURN( + Node * and_trail, + AndReduceTrailing(n->operand(0), trailing_ones, n->loc())); + std::vector args = {or_red, and_trail}; + XLS_RETURN_IF_ERROR( + n->ReplaceUsesWithNew(args, Op::kNor).status()); + return true; + } + case Op::kUGt: { + XLS_ASSIGN_OR_RETURN( + Node * or_red, + OrReduceLeading(n->operand(0), leading_zeros, n->loc())); + XLS_RETURN_IF_ERROR(n->ReplaceUsesWith(or_red)); + return true; + } + case Op::kULe: { + XLS_ASSIGN_OR_RETURN( + Node * nor_red, + NorReduceLeading(n->operand(0), leading_zeros, n->loc())); + XLS_RETURN_IF_ERROR(n->ReplaceUsesWith(nor_red)); + return true; + } + case Op::kUGe: { + XLS_ASSIGN_OR_RETURN( + Node * or_red, + OrReduceLeading(n->operand(0), leading_zeros, n->loc())); + XLS_ASSIGN_OR_RETURN( + Node * and_trail, + AndReduceTrailing(n->operand(0), trailing_ones, n->loc())); + std::vector args = {or_red, and_trail}; + XLS_RETURN_IF_ERROR( + n->ReplaceUsesWithNew(args, Op::kOr).status()); + return true; + } + default: + break; } - default: - break; } } diff --git a/xls/passes/arith_simplification_pass_test.cc b/xls/passes/arith_simplification_pass_test.cc index 16808a18bd..2f61292b6a 100644 --- a/xls/passes/arith_simplification_pass_test.cc +++ b/xls/passes/arith_simplification_pass_test.cc @@ -2165,6 +2165,23 @@ TEST_F(ArithSimplificationPassTest, UGeMask) { m::AndReduce(m::BitSlice(m::Param("x"), /*start=*/0, /*width=*/2)))); } +TEST_F(ArithSimplificationPassTest, UGtMaskMinusOne) { + auto p = CreatePackage(); + XLS_ASSERT_OK_AND_ASSIGN(Function * f, ParseFunction(R"( + fn f(x: bits[4]) -> bits[1] { + literal.1: bits[4] = literal(value=0b0010) + ret result: bits[1] = ugt(x, literal.1) + } +)", + p.get())); + ASSERT_THAT(Run(p.get()), IsOkAndHolds(true)); + EXPECT_THAT( + f->return_value(), + m::Or( + m::OrReduce(m::BitSlice(m::Param("x"), /*start=*/2, /*width=*/2)), + m::AndReduce(m::BitSlice(m::Param("x"), /*start=*/0, /*width=*/2)))); +} + TEST_F(ArithSimplificationPassTest, UGtMaskAllOnes) { auto p = CreatePackage(); XLS_ASSERT_OK_AND_ASSIGN(Function * f, ParseFunction(R"(