diff --git a/include/tc/core/cosets.hpp b/include/tc/core/cosets.hpp index 5afbcd0..d028019 100644 --- a/include/tc/core/cosets.hpp +++ b/include/tc/core/cosets.hpp @@ -29,21 +29,10 @@ struct Cosets { } } - void put(int idx, int target) { - // todo remove - int coset = idx / ngens; - int gen = idx % ngens; - put(coset, gen, target); - } - [[nodiscard]] int get(int coset, int gen) const { return data[coset * ngens + gen]; } - [[nodiscard]] int get(int idx) const { - return data[idx]; - } - [[nodiscard]] size_t size() const { return path.size(); } diff --git a/include/tc/core/solver.hpp b/include/tc/core/solver.hpp index 3d38b48..5470ae1 100644 --- a/include/tc/core/solver.hpp +++ b/include/tc/core/solver.hpp @@ -1,41 +1,110 @@ #pragma once +#include +#include + namespace { struct Row { int gnr; std::shared_ptr lst; }; - struct RelTables { + struct Table { private: + public: + Rel rel; std::vector rows; - size_t nrels; public: - explicit RelTables(size_t nrels) : nrels(nrels) { + explicit Table(const Rel &rel) : rel(rel) { } + }; - Row &operator()(size_t irel, size_t coset) { - // todo check bounds - size_t idx = coset * nrels + irel; - return rows[idx]; + struct Tables { + private: + std::shared_ptr null_lst_ptr = std::make_shared(); + int ngens; + + std::vector> tables; + std::vector>> deps; + + public: + explicit Tables(const tc::Group &group) + : ngens(group.ngens), deps(ngens) { + for (const auto &rel: group.get_rels()) { + auto table = std::make_shared(rel); + tables.push_back(table); + deps[rel.gens[0]].push_back(table); + deps[rel.gens[1]].push_back(table); + } } void add_row() { // std::vector already does block allocation. - rows.resize(rows.size() + nrels); + for (const auto &table: tables) { + table->rows.emplace_back(); + } + } + + void initialize(int target, const Cosets &cosets) { + for (auto &table: tables) { + const Rel &rel = table->rel; + Row &row = table->rows[target]; + + if (row.lst == nullptr) { + if (cosets.get(target, rel.gens[0]) != target and + cosets.get(target, rel.gens[1]) != target) { + row.lst = std::make_shared(); + row.gnr = 0; + } else { + row.lst = null_lst_ptr; + row.gnr = -1; + } + } + } + } + + void learn(int coset, int gen, int target, const Cosets &cosets, std::priority_queue &facts) { + if (target == coset) { + for (auto &table: deps[gen]) { + Row &target_row = table->rows[target]; + + if (target_row.lst == nullptr) { + target_row.gnr = -1; + } + } + } + + for (auto &table: deps[gen]) { + Row &target_row = table->rows[target]; + Row &coset_row = table->rows[coset]; + const Rel &rel = table->rel; + + if (target_row.lst == nullptr) { + target_row.lst = coset_row.lst; + target_row.gnr = coset_row.gnr + 1; + + if (coset_row.gnr < 0) { + target_row.gnr -= 2; + } + + if (target_row.gnr == rel.mult) { + // forward learn + int lst = *target_row.lst; + int gen_ = rel.gens[rel.gens[0] == gen]; + facts.push(lst * ngens + gen_); + } else if (target_row.gnr == -rel.mult) { + // stationary learn + int gen_ = rel.gens[rel.gens[0] == gen]; + facts.push(target * ngens + gen_); + } else if (target_row.gnr == rel.mult - 1) { + // determined family + *target_row.lst = target; + } + } + } } }; - - std::vector> dependency_map(int ngens, const std::vector &rels) { - std::vector> deps(ngens); - for (int irel = 0; irel < rels.size(); ++irel) { - const Rel &rel = rels[irel]; - deps[rel.gens[0]].push_back(irel); - deps[rel.gens[1]].push_back(irel); - } - return deps; - } } namespace tc { @@ -52,37 +121,21 @@ namespace tc { cosets.put(0, gen, 0); } - const auto &rels = g.get_rels(); // todo move to Group member - const auto nrels = rels.size(); - - auto deps = dependency_map(g.ngens, rels); - - std::shared_ptr null_lst_ptr = std::make_shared(); - - RelTables tables(nrels); + Tables tables(g); tables.add_row(); + tables.initialize(0, cosets); - for (int irel = 0; irel < nrels; irel++) { - Row &row = tables(irel, 0); - const Rel &rel = rels[irel]; - - if (cosets.get(rel.gens[0]) == -1 && cosets.get(rel.gens[1]) == -1) { - row.lst = std::make_shared(); - row.gnr = 0; - } else { - row.lst = null_lst_ptr; - row.gnr = -1; - } - } + std::priority_queue facts; for (int idx = 0; idx < cosets.data.size(); idx++) { - if (cosets.get(idx) >= 0) continue; + int coset = idx / g.ngens; + int gen = idx % g.ngens; + if (cosets.get(coset, gen) >= 0) continue; int target = cosets.size(); cosets.add_row(); tables.add_row(); - std::priority_queue facts; facts.push(idx); // todo nothing before the current coset will be used. @@ -94,68 +147,18 @@ namespace tc { int fact_idx = facts.top(); facts.pop(); - if (cosets.get(fact_idx) != -1) + int coset_ = fact_idx / g.ngens; + int gen_ = fact_idx % g.ngens; + + if (cosets.get(coset_, gen_) != -1) continue; - cosets.put(fact_idx, target); + cosets.put(coset_, gen_, target); - int coset = fact_idx / g.ngens; - int gen = fact_idx % g.ngens; - - if (target == coset) { - for (int irel: deps[gen]) { - Row &target_row = tables(irel, target); - if (target_row.lst == nullptr) { - target_row.gnr = -1; - } - } - } - - for (int irel: deps[gen]) { - Row &target_row = tables(irel, target); - Row &coset_row = tables(irel, coset); - - if (target_row.lst == nullptr) { - const Rel &rel = rels[irel]; - target_row.lst = coset_row.lst; - target_row.gnr = coset_row.gnr + 1; - - if (coset_row.gnr < 0) { - target_row.gnr -= 2; - } - - if (target_row.gnr == rel.mult) { - // forward learn - int lst = *target_row.lst; - int gen_ = rel.gens[rel.gens[0] == gen]; - facts.push(lst * g.ngens + gen_); - } else if (target_row.gnr == -rel.mult) { - // stationary learn - int gen_ = rel.gens[rel.gens[0] == gen]; - facts.push(target * g.ngens + gen_); - } else if (target_row.gnr == rel.mult - 1) { - // determined family - *target_row.lst = target; - } - } - } + tables.learn(coset_, gen_, target, cosets, facts); } - for (int irel = 0; irel < nrels; irel++) { - Row &target_row = tables(irel, target); - const Rel &rel = rels[irel]; - - if (target_row.lst == nullptr) { - if ((cosets.get(target, rel.gens[0]) != target) and - (cosets.get(target, rel.gens[1]) != target)) { - target_row.lst = std::make_shared(); - target_row.gnr = 0; - } else { - target_row.lst = null_lst_ptr; - target_row.gnr = -1; - } - } - } + tables.initialize(target, cosets); } return cosets;