diff --git a/src/ast/rewriter/term_enumeration.cpp b/src/ast/rewriter/term_enumeration.cpp index 2397db467b..eba649eacc 100644 --- a/src/ast/rewriter/term_enumeration.cpp +++ b/src/ast/rewriter/term_enumeration.cpp @@ -488,6 +488,106 @@ private: } }; +// ============================================================================ +// sort_stream - self-contained enumeration of terms of a single sort +// ============================================================================ + +/** + * A sort_stream owns a private grammar and bottom_up_enumerator so that + * several streams can be advanced independently and concurrently (as needed + * by the tuple iterator). It is seeded from the user-configured productions + * and, for array sorts, augments the grammar with fresh select operators and + * bound variables, wrapping enumerated bodies into lambdas. + */ +class sort_stream { +public: + sort_stream(ast_manager& m, func_decl_ref_vector const& funcs, expr_ref_vector const& exprs, sort* s) + : m(m), m_grammar(m), m_enum(m_grammar), autil(m), m_sort(s), m_current(m), m_pinned(m) { + for (func_decl* f : funcs) + m_grammar.add_func_decl(f); + for (expr* e : exprs) + m_grammar.add_expr(e); + m_enum.reset(); + init_sort(); + advance(); + } + + // Return the current term (already lambda-wrapped) and advance to the next + // one. Returns nullptr once the stream is exhausted. Returned terms remain + // pinned for the lifetime of the stream. + expr* next() { + if (m_end) + return nullptr; + expr* r = m_current.get(); + advance(); + return r; + } + +private: + ast_manager& m; + grammar m_grammar; + bottom_up_enumerator m_enum; + array_util autil; + sort* m_sort; + expr_ref m_current; + expr_ref_vector m_pinned; + bool m_end = false; + vector m_vars; + vector> m_decls; + vector> m_names; + + void init_sort() { + sort* range = m_sort; + while (autil.is_array(range)) { + m_vars.push_back(expr_ref_vector(m)); + m_decls.push_back(ptr_vector()); + m_names.push_back(vector()); + for (unsigned i = 0; i < get_array_arity(range); ++i) { + m_decls.back().push_back(get_array_domain(range, i)); + m_vars.back().push_back(nullptr); + m_names.back().push_back(symbol()); + } + expr_ref_vector args(m); + args.push_back(m.mk_const("a", range)); + for (unsigned i = 0; i < m_decls.back().size(); ++i) + args.push_back(m.mk_var(i, m_decls.back().get(i))); + app_ref sel(autil.mk_select(args), m); + m_grammar.add_func_decl(sel->get_decl()); + range = get_array_range(range); + } + unsigned n = 0; + for (unsigned i = m_decls.size(); i-- > 0;) { + for (unsigned j = m_decls[i].size(); j-- > 0;) { + m_vars[i][j] = m.mk_var(n, m_decls[i][j]); + m_names[i][j] = symbol(n); + m_grammar.add_expr(m_vars[i].get(j)); + n++; + } + } + m_sort = range; + m_enum.set_target_sort(range); + } + + void mk_lambda() { + if (!m_current) + return; + for (unsigned i = m_decls.size(); i-- > 0;) + m_current = m.mk_lambda(m_decls[i].size(), m_decls[i].data(), m_names[i].data(), m_current); + } + + void advance() { + if (m_end) + return; + m_current = m_enum.next(); + SASSERT(!m_current || m_current->get_sort() == m_sort); + mk_lambda(); + if (!m_current) + m_end = true; + else + m_pinned.push_back(m_current); + } +}; + } // namespace term_enum // ============================================================================ @@ -499,16 +599,21 @@ struct term_enumeration::imp { term_enum::grammar m_grammar; term_enum::bottom_up_enumerator m_bottom_up_enumerator; std::function m_cost; + func_decl_ref_vector m_user_funcs; + expr_ref_vector m_user_exprs; imp(ast_manager& m) : - m(m), m_grammar(m), m_bottom_up_enumerator(m_grammar) {} + m(m), m_grammar(m), m_bottom_up_enumerator(m_grammar), + m_user_funcs(m), m_user_exprs(m) {} void add_production(func_decl* f) { m_grammar.add_func_decl(f); + m_user_funcs.push_back(f); } void add_production(expr* e) { m_grammar.add_expr(e); + m_user_exprs.push_back(e); } void set_cost(std::function const& cost) { @@ -650,6 +755,157 @@ term_enumeration::iterator term_enumeration::terms::end() { return iterator(nullptr); } +// -- tuple iterator implementation -- + +struct term_enumeration::tuple_iterator::timp { + imp& m_imp; + ast_manager& m; + unsigned m_n; + scoped_ptr_vector m_streams; + vector m_terms; // materialized terms per dimension + svector m_done; // per-dimension exhausted flag + unsigned m_rr = 0; // round-robin cursor + vector m_buffer; // pending tuples + unsigned m_buf_idx = 0; + expr_ref_vector m_current; + bool m_end = false; + bool m_dead = false; + + timp(imp& i, unsigned n, sort* const* sorts) : + m_imp(i), m(i.m), m_n(n), m_current(i.m) { + for (unsigned k = 0; k < n; ++k) { + m_streams.push_back(alloc(term_enum::sort_stream, m, i.m_user_funcs, i.m_user_exprs, sorts[k])); + m_terms.push_back(expr_ref_vector(m)); + m_done.push_back(false); + } + if (n == 0) { + m_end = true; + return; + } + advance(); + } + + // Emit all tuples where dimension d is the fresh term t and every other + // dimension ranges over its currently materialized terms. + void emit(unsigned d, expr* t) { + for (unsigned j = 0; j < m_n; ++j) + if (j != d && m_terms[j].empty()) + return; + svector idx; + idx.resize(m_n, 0); + while (true) { + expr_ref_vector tup(m); + for (unsigned j = 0; j < m_n; ++j) + tup.push_back(j == d ? t : m_terms[j].get(idx[j])); + m_buffer.push_back(tup); + bool carried = false; + for (unsigned j = m_n; j-- > 0;) { + if (j == d) + continue; + idx[j]++; + if (idx[j] < m_terms[j].size()) { + carried = true; + break; + } + idx[j] = 0; + } + if (!carried) + break; + } + } + + // Advance one dimension. Returns false if every dimension is exhausted. + bool step() { + for (unsigned tries = 0; tries < m_n; ++tries) { + unsigned d = m_rr; + m_rr = (m_rr + 1) % m_n; + if (m_done[d]) + continue; + expr* t = m_streams[d]->next(); + if (!t) { + m_done[d] = true; + if (m_terms[d].empty()) + m_dead = true; // this dimension yields no terms at all + return true; + } + emit(d, t); + m_terms[d].push_back(t); + return true; + } + return false; + } + + bool fill() { + while (m_buf_idx >= m_buffer.size()) { + if (m_dead) + return false; + if (!step()) + return false; + } + return true; + } + + void advance() { + if (m_end) + return; + if (m_buf_idx >= m_buffer.size()) { + m_buffer.reset(); + m_buf_idx = 0; + if (!fill()) { + m_end = true; + m_current.reset(); + return; + } + } + m_current.reset(); + m_current.append(m_buffer[m_buf_idx]); + m_buf_idx++; + } +}; + +term_enumeration::tuple_iterator::tuple_iterator(imp& i, unsigned n, sort* const* sorts) { + m_imp = alloc(timp, i, n, sorts); +} + +term_enumeration::tuple_iterator::tuple_iterator(std::nullptr_t) { + m_imp = nullptr; +} + +term_enumeration::tuple_iterator::~tuple_iterator() { + dealloc(m_imp); +} + +expr_ref_vector term_enumeration::tuple_iterator::operator*() { + SASSERT(m_imp); + return m_imp->m_current; +} + +term_enumeration::tuple_iterator& term_enumeration::tuple_iterator::operator++() { + if (m_imp) + m_imp->advance(); + return *this; +} + +bool term_enumeration::tuple_iterator::operator==(tuple_iterator const& other) const { + if (!m_imp && !other.m_imp) return true; + if (!m_imp) return other.m_imp->m_end; + if (!other.m_imp) return m_imp->m_end; + return m_imp->m_end == other.m_imp->m_end; +} + +// -- tuples implementation -- + +term_enumeration::tuples::tuples(imp* i, unsigned n, sort* const* sorts) : + m_imp(i), m_sorts(n, sorts) {} + +term_enumeration::tuple_iterator term_enumeration::tuples::begin() { + return tuple_iterator(*m_imp, m_sorts.size(), m_sorts.data()); +} + +term_enumeration::tuple_iterator term_enumeration::tuples::end() { + return tuple_iterator(nullptr); +} + // -- term_enumeration implementation -- term_enumeration::term_enumeration(ast_manager& m) { @@ -676,6 +932,14 @@ term_enumeration::terms term_enumeration::enum_terms(sort* s) { return terms(m_imp, s); } +term_enumeration::tuples term_enumeration::enum_tuples(unsigned n, sort* const* sorts) { + return tuples(m_imp, n, sorts); +} + +term_enumeration::tuples term_enumeration::enum_tuples(sort_ref_vector const& sorts) { + return tuples(m_imp, sorts.size(), sorts.data()); +} + std::ostream& term_enumeration::display(std::ostream& out) const { return m_imp->display(out); } diff --git a/src/ast/rewriter/term_enumeration.h b/src/ast/rewriter/term_enumeration.h index 865b8d4021..5a77dfc03f 100644 --- a/src/ast/rewriter/term_enumeration.h +++ b/src/ast/rewriter/term_enumeration.h @@ -46,5 +46,36 @@ public: terms enum_terms(sort* s); + // -- tuple enumeration -- + // Iterate over vectors of terms, one term per input sort. Produces all + // combinations (dovetailed, since individual streams may be infinite). + + class tuple_iterator { + struct timp; + timp* m_imp; + public: + tuple_iterator(imp& i, unsigned n, sort* const* sorts); + tuple_iterator(std::nullptr_t); + ~tuple_iterator(); + expr_ref_vector operator*(); + tuple_iterator& operator++(); + bool operator!=(tuple_iterator const& other) const { + return !(*this == other); + } + bool operator==(tuple_iterator const& other) const; + }; + + class tuples { + imp* m_imp; + ptr_vector m_sorts; + public: + tuples(imp* i, unsigned n, sort* const* sorts); + tuple_iterator begin(); + tuple_iterator end(); + }; + + tuples enum_tuples(unsigned n, sort* const* sorts); + tuples enum_tuples(sort_ref_vector const& sorts); + std::ostream& display(std::ostream& out) const; }; \ No newline at end of file diff --git a/src/test/term_enumeration.cpp b/src/test/term_enumeration.cpp index 57b5da852f..20a87fb596 100644 --- a/src/test/term_enumeration.cpp +++ b/src/test/term_enumeration.cpp @@ -297,6 +297,63 @@ static void tst_nested_array_enumeration() { te.display(std::cout); } +static void tst_tuple_enumeration() { + std::cout << "=== test tuple enumeration ===\n"; + ast_manager m; + reg_decl_plugins(m); + arith_util a(m); + array_util arr(m); + + term_enumeration te(m); + + // Leaves: a boolean constant, integer constants, and an integer array. + expr_ref bt(m.mk_true(), m); + expr_ref bf(m.mk_false(), m); + expr_ref zero(a.mk_int(0), m); + expr_ref one(a.mk_int(1), m); + te.add_production(bt); + te.add_production(bf); + te.add_production(zero); + te.add_production(one); + + sort* int_sort = a.mk_int(); + sort_ref arr_sort(arr.mk_array_sort(int_sort, int_sort), m); + app_ref ca(arr.mk_const_array(arr_sort, zero), m); + te.add_production(ca.get()); + + // Operators over integers so that streams are non-trivial. + app_ref tmp_add(a.mk_add(zero, one), m); + te.add_production(tmp_add->get_decl()); + + sort_ref_vector sorts(m); + sorts.push_back(m.mk_bool_sort()); + sorts.push_back(int_sort); + sorts.push_back(arr_sort); + + unsigned count = 0; + obj_hashtable bools, ints, arrays; + for (expr_ref_vector const& tup : te.enum_tuples(sorts)) { + ENSURE(tup.size() == 3); + ENSURE(tup[0]->get_sort() == m.mk_bool_sort()); + ENSURE(tup[1]->get_sort() == int_sort); + ENSURE(tup[2]->get_sort() == arr_sort.get()); + bools.insert(tup[0]); + ints.insert(tup[1]); + arrays.insert(tup[2]); + if (count < 10) + std::cout << " (" << mk_pp(tup[0], m) << ", " << mk_pp(tup[1], m) + << ", " << mk_pp(tup[2], m) << ")\n"; + count++; + if (count >= 30) break; + } + + ENSURE(count >= 4); + // Verify combinations mix distinct terms from each component stream. + ENSURE(bools.size() >= 2); + ENSURE(ints.size() >= 2); + std::cout << "Enumerated " << count << " tuples\n"; +} + void tst_term_enumeration() { tst_basic_enumeration(); tst_enumeration_with_operators(); @@ -305,5 +362,6 @@ void tst_term_enumeration() { tst_bitvector_enumeration(); tst_multiple_sorts(); tst_nested_array_enumeration(); + tst_tuple_enumeration(); std::cout << "All term_enumeration tests passed!\n"; }