From 5b7023be051bfe0c803181a1043123fa31ef2ee2 Mon Sep 17 00:00:00 2001 From: Ryan James Date: Tue, 6 Oct 2026 15:10:07 -0600 Subject: [PATCH] Restructure: inference moves to sofic.inference; automata split into learning, vpa, enumeration, canonical packages - sofic.inference.hmm (filtering, em, information) replaces generators/hmm_inference.py; sampling moves to generators/sampling.py. - sofic.inference.cssr (counts, significance, process, subtree, transducer, stack) replaces generators/epsilon_inference.py, epsilon_transducer_inference.py, stack_inference.py and their private helpers; the spectral wrapper joins inference/spectral.py. - sofic.automata.learning (active, rpni, edsm, dfasat, alergia, papni, nlstar, observation), sofic.automata.vpa (base, operations, deterministic, modular, canonical, simulation), sofic.automata.enumeration (icdfa, idfa, words), and sofic.automata.canonical (residual, rfsa, atomaton, dual). - Old modules deleted without shims; docs pages and toctrees moved with them. - Doctests and prose updated for topological entropy in bits. Co-authored-by: Cursor --- docs/automata/atomaton.rst | 6 +- docs/automata/icdfa.rst | 10 +- docs/automata/learning.rst | 72 +- docs/automata/observation_table.rst | 10 +- docs/automata/rfsa.rst | 8 +- docs/automata/vpa.rst | 8 +- docs/generators/epsilon_machine.rst | 2 +- docs/generators/epsilon_transducer.rst | 2 +- docs/generators/generators.rst | 11 - .../cssr.rst} | 18 +- docs/inference/diagnostics.rst | 2 +- docs/inference/hdp_hmm.rst | 2 +- .../hmm_inference.rst => inference/hmm.rst} | 14 +- docs/inference/inference.rst | 11 +- docs/inference/model_selection.rst | 2 +- docs/inference/spectral.rst | 6 +- .../stack_cssr.rst} | 12 +- .../transcssr.rst} | 10 +- sofic/automata/__init__.py | 64 +- sofic/automata/base.py | 4 +- sofic/automata/canonical/__init__.py | 41 + sofic/automata/{ => canonical}/atomaton.py | 14 +- .../{canonical_dual.py => canonical/dual.py} | 8 +- .../residual.py} | 8 +- sofic/automata/{ => canonical}/rfsa.py | 12 +- sofic/automata/enumeration/__init__.py | 58 + sofic/automata/{ => enumeration}/icdfa.py | 0 sofic/automata/{ => enumeration}/idfa.py | 2 +- .../{enumeration.py => enumeration/words.py} | 0 sofic/automata/languages/atoms.py | 2 +- sofic/automata/languages/residuals.py | 2 +- sofic/automata/learning/__init__.py | 69 + sofic/automata/{ => learning}/active.py | 2 +- sofic/automata/{ => learning}/alergia.py | 2 +- sofic/automata/{ => learning}/dfasat.py | 6 +- sofic/automata/{ => learning}/edsm.py | 2 +- .../{learning.py => learning/nlstar.py} | 10 +- sofic/automata/{ => learning}/observation.py | 12 +- sofic/automata/{ => learning}/papni.py | 2 +- sofic/automata/{ => learning}/rpni.py | 0 sofic/automata/vpa.py | 1331 ----------------- sofic/automata/vpa/__init__.py | 26 + sofic/automata/vpa/base.py | 251 ++++ sofic/automata/vpa/canonical.py | 280 ++++ sofic/automata/vpa/deterministic.py | 166 ++ sofic/automata/vpa/modular.py | 686 +++++++++ .../operations.py} | 15 +- .../{vpa_simulation.py => vpa/simulation.py} | 6 +- sofic/generators/__init__.py | 15 - sofic/generators/_suffix_counts.py | 21 - sofic/generators/base.py | 24 +- sofic/generators/epsilon_machine.py | 12 +- sofic/generators/epsilon_transducer.py | 2 +- sofic/generators/hmm_inference.py | 688 --------- sofic/generators/sampling.py | 45 + .../topological_epsilon_enumeration.py | 4 +- sofic/inference/__init__.py | 54 +- sofic/inference/cssr/__init__.py | 54 + sofic/inference/cssr/counts.py | 249 +++ .../cssr/process.py} | 366 +---- .../cssr/significance.py} | 154 +- .../cssr/stack.py} | 86 +- sofic/inference/cssr/subtree.py | 125 ++ .../cssr/transducer.py} | 145 +- sofic/inference/hmm/__init__.py | 37 + sofic/inference/hmm/em.py | 239 +++ sofic/inference/hmm/filtering.py | 237 +++ sofic/inference/hmm/information.py | 208 +++ sofic/inference/model_selection.py | 6 +- sofic/inference/spectral.py | 64 +- sofic/serialization.py | 4 +- sofic/testing/strategies.py | 2 +- tests/test_active_learning.py | 6 +- tests/test_atomaton.py | 4 +- tests/test_canonical_rfsa.py | 6 +- tests/test_epsilon_inference.py | 11 +- tests/test_epsilon_transducer_inference.py | 4 +- tests/test_hmm_inference.py | 12 +- tests/test_icdfa.py | 2 +- tests/test_inference_diagnostics.py | 4 +- tests/test_learning.py | 8 +- tests/test_observation.py | 2 +- tests/test_rfsa.py | 4 +- tests/test_stack_inference.py | 16 +- tests/test_testing_strategies.py | 2 +- tests/test_topological_epsilon.py | 4 +- tests/test_vpa_constructions.py | 2 +- tests/test_yaml.py | 2 +- 88 files changed, 3325 insertions(+), 2892 deletions(-) rename docs/{generators/epsilon_inference.rst => inference/cssr.rst} (94%) rename docs/{generators/hmm_inference.rst => inference/hmm.rst} (89%) rename docs/{generators/stack_inference.rst => inference/stack_cssr.rst} (87%) rename docs/{generators/epsilon_transducer_inference.rst => inference/transcssr.rst} (84%) create mode 100644 sofic/automata/canonical/__init__.py rename sofic/automata/{ => canonical}/atomaton.py (88%) rename sofic/automata/{canonical_dual.py => canonical/dual.py} (76%) rename sofic/automata/{canonical_extraction.py => canonical/residual.py} (97%) rename sofic/automata/{ => canonical}/rfsa.py (80%) create mode 100644 sofic/automata/enumeration/__init__.py rename sofic/automata/{ => enumeration}/icdfa.py (100%) rename sofic/automata/{ => enumeration}/idfa.py (99%) rename sofic/automata/{enumeration.py => enumeration/words.py} (100%) create mode 100644 sofic/automata/learning/__init__.py rename sofic/automata/{ => learning}/active.py (99%) rename sofic/automata/{ => learning}/alergia.py (99%) rename sofic/automata/{ => learning}/dfasat.py (95%) rename sofic/automata/{ => learning}/edsm.py (99%) rename sofic/automata/{learning.py => learning/nlstar.py} (96%) rename sofic/automata/{ => learning}/observation.py (68%) rename sofic/automata/{ => learning}/papni.py (99%) rename sofic/automata/{ => learning}/rpni.py (100%) delete mode 100644 sofic/automata/vpa.py create mode 100644 sofic/automata/vpa/__init__.py create mode 100644 sofic/automata/vpa/base.py create mode 100644 sofic/automata/vpa/canonical.py create mode 100644 sofic/automata/vpa/deterministic.py create mode 100644 sofic/automata/vpa/modular.py rename sofic/automata/{vpa_constructions.py => vpa/operations.py} (98%) rename sofic/automata/{vpa_simulation.py => vpa/simulation.py} (87%) delete mode 100644 sofic/generators/_suffix_counts.py delete mode 100644 sofic/generators/hmm_inference.py create mode 100644 sofic/generators/sampling.py create mode 100644 sofic/inference/cssr/__init__.py create mode 100644 sofic/inference/cssr/counts.py rename sofic/{generators/epsilon_inference.py => inference/cssr/process.py} (55%) rename sofic/{generators/_morph_tests.py => inference/cssr/significance.py} (54%) rename sofic/{generators/stack_inference.py => inference/cssr/stack.py} (87%) create mode 100644 sofic/inference/cssr/subtree.py rename sofic/{generators/epsilon_transducer_inference.py => inference/cssr/transducer.py} (73%) create mode 100644 sofic/inference/hmm/__init__.py create mode 100644 sofic/inference/hmm/em.py create mode 100644 sofic/inference/hmm/filtering.py create mode 100644 sofic/inference/hmm/information.py diff --git a/docs/automata/atomaton.rst b/docs/automata/atomaton.rst index 1f5b3b0..12ba555 100644 --- a/docs/automata/atomaton.rst +++ b/docs/automata/atomaton.rst @@ -1,5 +1,5 @@ .. atomaton.rst -.. py:module:: sofic.automata.atomaton +.. py:module:: sofic.automata.canonical.atomaton ******** Átomaton @@ -29,7 +29,7 @@ the special case where the reverse is deterministic. .. code-block:: python - from sofic.automata.atomaton import Atomaton, atomic_states, is_atomic + from sofic.automata.canonical.atomaton import Atomaton, atomic_states, is_atomic atomaton = Atomaton.from_language(dfa) is_atomic(atomaton) # True @@ -58,4 +58,4 @@ API .. autoclass:: MaximizedPrimeAtomaton :members: from_language, from_canonical_rfsa, dual -.. autofunction:: sofic.automata.canonical_extraction.maximized_prime_atomaton_from_language +.. autofunction:: sofic.automata.canonical.residual.maximized_prime_atomaton_from_language diff --git a/docs/automata/icdfa.rst b/docs/automata/icdfa.rst index d4c0f95..5e905ab 100644 --- a/docs/automata/icdfa.rst +++ b/docs/automata/icdfa.rst @@ -1,5 +1,5 @@ .. icdfa.rst -.. py:module:: sofic.automata.icdfa +.. py:module:: sofic.automata.enumeration.icdfa ***** ICDFA @@ -23,7 +23,7 @@ API .. autofunction:: count_icdfa .. autofunction:: count_icdfa_empty -.. autofunction:: sofic.automata.idfa.iter_idfa_strings -.. autofunction:: sofic.automata.idfa.rank_idfa_string -.. autofunction:: sofic.automata.idfa.unrank_idfa_string -.. autofunction:: sofic.automata.idfa.count_accessible_idfa +.. autofunction:: sofic.automata.enumeration.idfa.iter_idfa_strings +.. autofunction:: sofic.automata.enumeration.idfa.rank_idfa_string +.. autofunction:: sofic.automata.enumeration.idfa.unrank_idfa_string +.. autofunction:: sofic.automata.enumeration.idfa.count_accessible_idfa diff --git a/docs/automata/learning.rst b/docs/automata/learning.rst index 18b879c..e8ab557 100644 --- a/docs/automata/learning.rst +++ b/docs/automata/learning.rst @@ -1,5 +1,5 @@ .. learning.rst -.. py:module:: sofic.automata.learning +.. py:module:: sofic.automata.learning.nlstar ******** Learning @@ -18,25 +18,25 @@ hypothesis is the canonical RFSA of the target (:doc:`rfsa`). Running NL\* on the reversed target and reversing the result learns the maximized prime átomaton (:doc:`atomaton`). -:class:`~sofic.automata.active.AutomatonEquivalenceOracle` answers equivalence +:class:`~sofic.automata.learning.active.AutomatonEquivalenceOracle` answers equivalence queries exactly against a target automaton, returning a shortest counterexample. -.. autofunction:: sofic.automata.learning.learn_rfsa_nlstar -.. autofunction:: sofic.automata.learning.learn_prime_atomaton_nlstar -.. autofunction:: sofic.automata.learning.learn_rfsa_from_language +.. autofunction:: sofic.automata.learning.nlstar.learn_rfsa_nlstar +.. autofunction:: sofic.automata.learning.nlstar.learn_prime_atomaton_nlstar +.. autofunction:: sofic.automata.learning.nlstar.learn_rfsa_from_language Active learning (L\*, TTT, Mealy) ================================= -:mod:`sofic.automata.active` learns an automaton from a *teacher* answering +:mod:`sofic.automata.learning.active` learns an automaton from a *teacher* answering **membership** and **equivalence** queries. It provides Angluin's L\* :cite:`Angluin1987` and a redundancy-free **discrimination-tree** learner in the TTT family :cite:`KearnsVazirani1994,Isberner2014` for :class:`~sofic.automata.dfa.DFA`, plus the Mealy variant of L\* :cite:`Shahbaz2009`; all use Rivest-Schapire counterexample analysis :cite:`RivestSchapire1993`. Oracles adapt sofic models: a -:class:`~sofic.automata.active.LanguageMembershipOracle` wraps any model exposing +:class:`~sofic.automata.learning.active.LanguageMembershipOracle` wraps any model exposing ``recognizes`` / ``__contains__`` (a :class:`~sofic.automata.dfa.DFA`, NFA, átomaton, or ``model.to_support_dfa()`` for a sofic shift or ε-machine), and the equivalence oracles offer bounded-exhaustive or random-walk testing. @@ -53,7 +53,7 @@ equivalence test: .. code-block:: python - from sofic.automata.active import ( + from sofic.automata.learning.active import ( FunctionMembershipOracle, RandomWalkEquivalenceOracle, learn_dfa_lstar, @@ -63,23 +63,23 @@ equivalence test: equivalence = RandomWalkEquivalenceOracle(membership, {"a", "b"}, rng=0) dfa = learn_dfa_lstar({"a", "b"}, membership, equivalence) -.. autofunction:: sofic.automata.active.learn_dfa_lstar -.. autofunction:: sofic.automata.active.learn_dfa_ttt -.. autofunction:: sofic.automata.active.learn_mealy_lstar -.. autofunction:: sofic.automata.active.learn_dfa_from_language -.. autofunction:: sofic.automata.active.learn_mealy_from_transducer +.. autofunction:: sofic.automata.learning.active.learn_dfa_lstar +.. autofunction:: sofic.automata.learning.active.learn_dfa_ttt +.. autofunction:: sofic.automata.learning.active.learn_mealy_lstar +.. autofunction:: sofic.automata.learning.active.learn_dfa_from_language +.. autofunction:: sofic.automata.learning.active.learn_mealy_from_transducer -.. autoclass:: sofic.automata.active.MembershipOracle +.. autoclass:: sofic.automata.learning.active.MembershipOracle :members: -.. autoclass:: sofic.automata.active.EquivalenceOracle +.. autoclass:: sofic.automata.learning.active.EquivalenceOracle :members: -.. autoclass:: sofic.automata.active.LanguageMembershipOracle -.. autoclass:: sofic.automata.active.FunctionMembershipOracle -.. autoclass:: sofic.automata.active.AutomatonEquivalenceOracle -.. autoclass:: sofic.automata.active.ExhaustiveEquivalenceOracle -.. autoclass:: sofic.automata.active.RandomWalkEquivalenceOracle -.. autoclass:: sofic.automata.active.TransducerOutputOracle -.. autoclass:: sofic.automata.active.MealyExhaustiveEquivalenceOracle +.. autoclass:: sofic.automata.learning.active.LanguageMembershipOracle +.. autoclass:: sofic.automata.learning.active.FunctionMembershipOracle +.. autoclass:: sofic.automata.learning.active.AutomatonEquivalenceOracle +.. autoclass:: sofic.automata.learning.active.ExhaustiveEquivalenceOracle +.. autoclass:: sofic.automata.learning.active.RandomWalkEquivalenceOracle +.. autoclass:: sofic.automata.learning.active.TransducerOutputOracle +.. autoclass:: sofic.automata.learning.active.MealyExhaustiveEquivalenceOracle Passive learning (RPNI) ======================= @@ -95,7 +95,7 @@ consistent with the sample :cite:`Lang1998`: dfa = learn_dfa_rpni(positive=["ab", "abab"], negative=["a", "b"]) dfa.validate() -.. autofunction:: sofic.automata.rpni.learn_dfa_rpni +.. autofunction:: sofic.automata.learning.rpni.learn_dfa_rpni Passive learning (EDSM / blue-fringe) ===================================== @@ -116,12 +116,12 @@ automaton: dfa = learn_dfa_edsm(positive=["a", "aba", "ababa"], negative=["", "b", "ab"]) dfa.validate() -.. autofunction:: sofic.automata.edsm.learn_dfa_edsm +.. autofunction:: sofic.automata.learning.edsm.learn_dfa_edsm Exact minimal DFA (SAT) ======================= -Where RPNI and EDSM are heuristics, :func:`sofic.automata.dfasat.learn_dfa_sat` +Where RPNI and EDSM are heuristics, :func:`sofic.automata.learning.dfasat.learn_dfa_sat` returns the **provably minimal** DFA consistent with the sample. Following Heule & Verwer :cite:`HeuleVerwer2010`, it translates the augmented prefix-tree acceptor into a graph-colouring SAT instance and searches the state count ``k`` @@ -136,7 +136,7 @@ satisfiable ``k``. It requires the optional `python-sat dfa = learn_dfa_sat(positive=["a", "aba", "ababa"], negative=["", "b", "ab"]) dfa.validate() -.. autofunction:: sofic.automata.dfasat.learn_dfa_sat +.. autofunction:: sofic.automata.learning.dfasat.learn_dfa_sat Probabilistic passive learning (ALERGIA) ======================================== @@ -146,7 +146,7 @@ from **unlabeled** positive strings by merging states of a frequency prefix-tree acceptor whenever a Hoeffding-bound test cannot distinguish their transition statistics :cite:`Carrasco1994`. It is the stochastic, unlabeled counterpart of RPNI/EDSM and a state-merging alternative to CSSR -(:func:`sofic.generators.epsilon_inference.cssr`). The compatibility threshold +(:func:`sofic.inference.cssr.process.cssr`). The compatibility threshold ``alpha`` trades off model size against fidelity: smaller ``alpha`` merges more aggressively (fewer states); larger ``alpha`` is more conservative. @@ -159,13 +159,13 @@ aggressively (fewer states); larger ``alpha`` is more conservative. pfa = learn_pfa_alergia(samples, alpha=0.05) pfa.validate() -.. autofunction:: sofic.automata.alergia.learn_pfa_alergia +.. autofunction:: sofic.automata.learning.alergia.learn_pfa_alergia Passive learning (PAPNI) ======================== PAPNI extends passive inference to visibly pushdown languages. Words over a -:class:`~sofic.automata.papni.DyckAlphabet` are stack-encoded, a DFA is +:class:`~sofic.automata.learning.papni.DyckAlphabet` are stack-encoded, a DFA is induced over the encoding, and the result is decoded to a :class:`~sofic.shifts.sofic_dyck.SoficDyckShift` :cite:`Muskardin2025`: @@ -181,13 +181,13 @@ induced over the encoding, and the result is decoded to a shift = learn_sofic_dyck_shift_papni(positive=["()", "(())"], negative=["("], alphabet=alphabet) For fitting probabilities on the learned topology, see -:doc:`../generators/stack_inference`. +:doc:`../inference/stack_cssr`. -.. autoclass:: sofic.automata.papni.DyckAlphabet +.. autoclass:: sofic.automata.learning.papni.DyckAlphabet :members: classify, symbol_alphabet -.. autofunction:: sofic.automata.papni.learn_sofic_dyck_shift_papni -.. autofunction:: sofic.automata.papni.is_well_matched -.. autofunction:: sofic.automata.papni.papni_encode -.. autofunction:: sofic.automata.papni.papni_encode_samples -.. autofunction:: sofic.automata.papni.sofic_dyck_shift_from_papni_dfa +.. autofunction:: sofic.automata.learning.papni.learn_sofic_dyck_shift_papni +.. autofunction:: sofic.automata.learning.papni.is_well_matched +.. autofunction:: sofic.automata.learning.papni.papni_encode +.. autofunction:: sofic.automata.learning.papni.papni_encode_samples +.. autofunction:: sofic.automata.learning.papni.sofic_dyck_shift_from_papni_dfa diff --git a/docs/automata/observation_table.rst b/docs/automata/observation_table.rst index 0e7deed..ea2b29c 100644 --- a/docs/automata/observation_table.rst +++ b/docs/automata/observation_table.rst @@ -1,5 +1,5 @@ .. observation_table.rst -.. py:module:: sofic.automata.observation +.. py:module:: sofic.automata.learning.observation ***************** Observation Table @@ -15,7 +15,7 @@ API .. autoclass:: ObservationTable -.. autofunction:: sofic.automata.canonical_extraction.observation_to_canonical_rfsa -.. autofunction:: sofic.automata.canonical_extraction.observation_to_atomaton -.. autofunction:: sofic.automata.canonical_extraction.observation_to_maximized_prime_atomaton -.. autofunction:: sofic.automata.canonical_extraction.observation_to_minimal_dfa +.. autofunction:: sofic.automata.canonical.residual.observation_to_canonical_rfsa +.. autofunction:: sofic.automata.canonical.residual.observation_to_atomaton +.. autofunction:: sofic.automata.canonical.residual.observation_to_maximized_prime_atomaton +.. autofunction:: sofic.automata.canonical.residual.observation_to_minimal_dfa diff --git a/docs/automata/rfsa.rst b/docs/automata/rfsa.rst index bd21737..3f23c22 100644 --- a/docs/automata/rfsa.rst +++ b/docs/automata/rfsa.rst @@ -1,5 +1,5 @@ .. rfsa.rst -.. py:module:: sofic.automata.rfsa +.. py:module:: sofic.automata.canonical.rfsa *** RFSA @@ -25,7 +25,7 @@ canonical RFSA from queries (:doc:`learning`). .. code-block:: python - from sofic.automata.rfsa import CanonicalRFSA + from sofic.automata.canonical.rfsa import CanonicalRFSA rfsa = CanonicalRFSA.from_language(nfa) rfsa.validate() # every state accepts a residual @@ -38,6 +38,6 @@ API .. autoclass:: CanonicalRFSA :members: from_language, from_observation_table, dual -.. autofunction:: sofic.automata.canonical_extraction.canonical_rfsa_from_language -.. autoclass:: sofic.automata.canonical_extraction.ResidualTable +.. autofunction:: sofic.automata.canonical.residual.canonical_rfsa_from_language +.. autoclass:: sofic.automata.canonical.residual.ResidualTable :members: includes, is_covered, prime_states diff --git a/docs/automata/vpa.rst b/docs/automata/vpa.rst index ae6c36d..22fd902 100644 --- a/docs/automata/vpa.rst +++ b/docs/automata/vpa.rst @@ -74,8 +74,8 @@ forms exist for well-matched languages, or once calls are assigned to modules The shared modular generalization: a call's target depends only on the call symbol. -:func:`~sofic.automata.vpa_constructions.to_single_entry` and -:func:`~sofic.automata.vpa_constructions.to_multiple_entry` convert any VPA of a +:func:`~sofic.automata.vpa.operations.to_single_entry` and +:func:`~sofic.automata.vpa.operations.to_multiple_entry` convert any VPA of a well-matched language into these forms (default: one module per call symbol). On a call the state resets to the module's entry and the caller is pushed, so the automaton forgets its caller; that is why pending calls -- and hence @@ -114,8 +114,8 @@ API .. autoclass:: CanonicalVisiblyPushdownAutomaton :members: from_vpa -.. automodule:: sofic.automata.vpa_constructions +.. automodule:: sofic.automata.vpa.operations :members: normalize, determinize, complement, concat, kleene_star, well_matched_summaries, accepted_word, has_unmatched_word, to_single_entry, to_multiple_entry -.. autofunction:: sofic.automata.vpa_simulation.recognizes_vpa +.. autofunction:: sofic.automata.vpa.simulation.recognizes_vpa diff --git a/docs/generators/epsilon_machine.rst b/docs/generators/epsilon_machine.rst index dc7d3af..e6e3264 100644 --- a/docs/generators/epsilon_machine.rst +++ b/docs/generators/epsilon_machine.rst @@ -87,7 +87,7 @@ structure can differ, because a cycle weight may equal one only by algebraic coincidence in the probabilities. See also :doc:`bidirectional_epsilon_machine`, :doc:`information_anatomy`, -:doc:`block_convergence`, and :doc:`epsilon_inference` (sample-based reconstruction). +:doc:`block_convergence`, and :doc:`../inference/cssr` (sample-based reconstruction). API === diff --git a/docs/generators/epsilon_transducer.rst b/docs/generators/epsilon_transducer.rst index 3ff0ffe..38b29b6 100644 --- a/docs/generators/epsilon_transducer.rst +++ b/docs/generators/epsilon_transducer.rst @@ -41,7 +41,7 @@ Minimize a joint-unifilar stochastic transducer to its causal states: The channel can also be read off a joint ``(input, output)`` generator with :meth:`~EpsilonTransducer.from_joint_generator`, or reconstructed from paired sample sequences with :meth:`~EpsilonTransducer.from_paired_sequences` (the -transCSSR algorithm, :doc:`epsilon_transducer_inference`). +transCSSR algorithm, :doc:`../inference/transcssr`). Channel measures ================ diff --git a/docs/generators/generators.rst b/docs/generators/generators.rst index 7f66ba2..87b695b 100644 --- a/docs/generators/generators.rst +++ b/docs/generators/generators.rst @@ -50,17 +50,6 @@ Constructions and conversions conversions lumping -Inference -========= - -.. toctree:: - :maxdepth: 1 - - hmm_inference - epsilon_inference - epsilon_transducer_inference - stack_inference - Advanced ======== diff --git a/docs/generators/epsilon_inference.rst b/docs/inference/cssr.rst similarity index 94% rename from docs/generators/epsilon_inference.rst rename to docs/inference/cssr.rst index f4808b3..a493bc0 100644 --- a/docs/generators/epsilon_inference.rst +++ b/docs/inference/cssr.rst @@ -1,5 +1,5 @@ -.. epsilon_inference.rst -.. py:module:: sofic.generators.epsilon_inference +.. cssr.rst +.. py:module:: sofic.inference.cssr ************************** ε-Machine inference @@ -99,7 +99,7 @@ apply directly. After reconstruction, check the result with :func:`~sofic.inference.diagnostics.goodness_of_fit` and :func:`~sofic.inference.diagnostics.structure_stability` -(see :doc:`../inference/diagnostics`). +(see :doc:`diagnostics`). .. autofunction:: morphs_differ @@ -139,22 +139,22 @@ process). Otherwise the Hankel singular-value gap selects the rank. observations, method="spectral", prefix_length=3, rank=2 ) -.. autofunction:: spectral +.. autofunction:: sofic.inference.spectral.spectral Unified entry point =================== Use :meth:`~sofic.generators.epsilon_machine.EpsilonMachine.from_sequence` to dispatch to CSSR, subtree merging, or spectral reconstruction -(see :doc:`epsilon_machine`). +(see :doc:`../generators/epsilon_machine`). Related inference methods ========================= -**transCSSR** — input/output ε-transducers; see :doc:`epsilon_transducer_inference`. +**transCSSR** — input/output ε-transducers; see :doc:`transcssr`. **Bayesian structural inference** — conjugate Dirichlet–multinomial evidence over -candidate unifilar topologies :cite:`Strelioff2014`; see :doc:`../inference/epsilon`. +candidate unifilar topologies :cite:`Strelioff2014`; see :doc:`epsilon`. This is not a Gibbs clustering heuristic over history labels. Subtree merging is the literature reconstruction by morph clustering; a separate @@ -173,5 +173,5 @@ the causal-state construction used after spectral learning. **RKHS ε-machines** — continuous-time extension (arXiv:2011.14821). -See also :doc:`epsilon_machine`, :doc:`constructions`, :doc:`hmm_inference`, -and :doc:`../inference/spectral`. +See also :doc:`../generators/epsilon_machine`, :doc:`../generators/constructions`, :doc:`hmm`, +and :doc:`spectral`. diff --git a/docs/inference/diagnostics.rst b/docs/inference/diagnostics.rst index a674dc9..a6ec334 100644 --- a/docs/inference/diagnostics.rst +++ b/docs/inference/diagnostics.rst @@ -29,7 +29,7 @@ usually means ``Lmax`` is shorter than the source's synchronization length. .. code-block:: python - from sofic.generators.epsilon_inference import cssr + from sofic.inference.cssr import cssr from sofic.inference.diagnostics import goodness_of_fit machine = cssr(data, Lmax=1) diff --git a/docs/inference/hdp_hmm.rst b/docs/inference/hdp_hmm.rst index 6a6ae46..400083b 100644 --- a/docs/inference/hdp_hmm.rst +++ b/docs/inference/hdp_hmm.rst @@ -60,7 +60,7 @@ The HDP-HMM infers the state count nonparametrically, whereas :doc:`model_selection` scores a *fixed* set of candidate orders with information criteria and :mod:`sofic.inference.bayesian` compares fixed unifilar topologies by exact Dirichlet-multinomial evidence. For point-estimate reconstruction see -CSSR in :doc:`../generators/epsilon_inference`. +CSSR in :doc:`cssr`. API === diff --git a/docs/generators/hmm_inference.rst b/docs/inference/hmm.rst similarity index 89% rename from docs/generators/hmm_inference.rst rename to docs/inference/hmm.rst index f4ae796..d51819c 100644 --- a/docs/generators/hmm_inference.rst +++ b/docs/inference/hmm.rst @@ -1,5 +1,5 @@ -.. hmm_inference.rst -.. py:module:: sofic.generators.hmm_inference +.. hmm.rst +.. py:module:: sofic.inference.hmm ************* HMM Inference @@ -19,7 +19,7 @@ via the Fisher and Louis :cite:`Louis1982` identities. In [3]: obs = [0, 1, 0, 0, 1] - In [4]: from sofic.generators.hmm_inference import forward, viterbi, sample + In [4]: from sofic.generators.sampling import sample; from sofic.inference.hmm import forward, viterbi In [5]: alpha = forward(eps, obs) @@ -33,7 +33,7 @@ pairs. .. ipython:: - In [8]: from sofic.generators.hmm_inference import smooth, two_slice_marginals + In [8]: from sofic.inference.hmm import smooth, two_slice_marginals In [9]: gamma = smooth(eps, obs) # gamma[t, s] = P(X_t = s | Y) @@ -44,7 +44,7 @@ transition-graph topology fixed: .. ipython:: - In [11]: from sofic.generators.hmm_inference import baum_welch + In [11]: from sofic.inference.hmm import baum_welch In [12]: data, _states = sample(golden_mean(0.3), n=2000) @@ -63,7 +63,7 @@ parameter uncertainty at the current parameters: .. ipython:: - In [14]: from sofic.generators.hmm_inference import score, standard_errors + In [14]: from sofic.inference.hmm import score, standard_errors In [15]: g = score(eps, obs) @@ -79,7 +79,7 @@ Filtering, decoding, and sampling .. autofunction:: backward .. autofunction:: log_likelihood .. autofunction:: viterbi -.. autofunction:: sample +.. autofunction:: sofic.generators.sampling.sample Smoothing --------- diff --git a/docs/inference/inference.rst b/docs/inference/inference.rst index e91692c..b0105a0 100644 --- a/docs/inference/inference.rst +++ b/docs/inference/inference.rst @@ -33,13 +33,18 @@ The historical names ``InferMC`` and ``InferEM`` are retained as aliases for For *non-Bayesian* reconstruction — Causal-State Splitting Reconstruction (CSSR), subtree merging, and spectral mixed-state extraction — see the point-estimate routines in - :doc:`../generators/epsilon_inference`, - :doc:`../generators/hmm_inference`, and - :doc:`../generators/stack_inference`. + :doc:`cssr`, + :doc:`transcssr`, + :doc:`hmm`, and + :doc:`stack_cssr`. .. toctree:: :maxdepth: 1 + hmm + cssr + transcssr + stack_cssr markov epsilon spectral diff --git a/docs/inference/model_selection.rst b/docs/inference/model_selection.rst index 6fefa30..1271926 100644 --- a/docs/inference/model_selection.rst +++ b/docs/inference/model_selection.rst @@ -9,7 +9,7 @@ Classical model-selection criteria for choosing the order (state count) of a fitted generator when a fully-Bayesian evidence is unavailable or undesirable. The module scores any fitted :class:`~sofic.generators.base.HiddenMarkovModel` using the log-likelihood (in bits) from -:func:`sofic.generators.hmm_inference.log_likelihood` and a free-parameter count +:func:`sofic.inference.hmm.log_likelihood` and a free-parameter count read off the transition graph. Log-likelihoods, log scores, and the MDL code length are in bits; AIC, AICc, BIC, and WAIC are computed from the natural log-likelihood so they keep their standard deviance scale: diff --git a/docs/inference/spectral.rst b/docs/inference/spectral.rst index 478542f..8f84b0b 100644 --- a/docs/inference/spectral.rst +++ b/docs/inference/spectral.rst @@ -12,7 +12,7 @@ matrix** of block probabilities, takes its singular value decomposition, and reads the observable operators off the truncated factorization :cite:`Balle2014`. It is the automata-theoretic twin of the spectral hidden Markov model algorithm of Hsu, Kakade & Zhang :cite:`Hsu2012`. Unlike -Baum-Welch (:func:`sofic.generators.hmm_inference.baum_welch`), spectral +Baum-Welch (:func:`sofic.inference.hmm.baum_welch`), spectral learning is a consistent, one-shot estimator with no local optima, and the model order is read from the singular-value spectrum instead of being fixed in advance. @@ -68,12 +68,12 @@ Projection to an ε-machine non-negative Mealy projection when one exists in the learned basis, otherwise mixed-state enumeration of the observable operators :cite:`Ellison2009`. The same path is -:func:`~sofic.generators.epsilon_inference.spectral` / +:func:`~sofic.inference.spectral.spectral` / ``EpsilonMachine.from_sequence(..., method="spectral")``. .. code-block:: python - from sofic.generators.epsilon_inference import spectral + from sofic.inference.spectral import spectral from sofic.examples import golden_mean process = golden_mean(0.5) diff --git a/docs/generators/stack_inference.rst b/docs/inference/stack_cssr.rst similarity index 87% rename from docs/generators/stack_inference.rst rename to docs/inference/stack_cssr.rst index af85695..0b2ab14 100644 --- a/docs/generators/stack_inference.rst +++ b/docs/inference/stack_cssr.rst @@ -1,5 +1,5 @@ -.. stack_inference.rst -.. py:module:: sofic.generators.stack_inference +.. stack_cssr.rst +.. py:module:: sofic.inference.cssr.stack ********************* Stack-HMM Inference @@ -7,8 +7,8 @@ Stack-HMM Inference Inference routines that reconstruct a :class:`~sofic.generators.stack_hmm.HiddenMarkovStackModel` from sequences -over a :class:`~sofic.automata.papni.DyckAlphabet`. These extend the -finite-state inference of :doc:`epsilon_inference` with visibly pushdown stack +over a :class:`~sofic.automata.learning.papni.DyckAlphabet`. These extend the +finite-state inference of :doc:`cssr` with visibly pushdown stack semantics :cite:`BealBlockeletDima2015`. Two families are provided: @@ -28,7 +28,7 @@ Two families are provided: .. code-block:: python from sofic.automata import DyckAlphabet - from sofic.generators import stack_cssr + from sofic.inference.cssr import stack_cssr alphabet = DyckAlphabet( call_alphabet=frozenset({"("}), @@ -49,7 +49,7 @@ the finite control. Return edges are matched only to calls observed to close them. ``stack_cssr`` accepts the same calibration options as -:func:`~sofic.generators.epsilon_inference.cssr`: ``test="exact"``, +:func:`~sofic.inference.cssr.cssr`: ``test="exact"``, ``correction="bonferroni"`` (over eligible configurations), and ``Lmax="auto"``. Stack processes generally have infinite Markov order, so the automatic depth is a lower bound on the suffix length the data support. diff --git a/docs/generators/epsilon_transducer_inference.rst b/docs/inference/transcssr.rst similarity index 84% rename from docs/generators/epsilon_transducer_inference.rst rename to docs/inference/transcssr.rst index da3280e..dc09af1 100644 --- a/docs/generators/epsilon_transducer_inference.rst +++ b/docs/inference/transcssr.rst @@ -1,11 +1,11 @@ -.. epsilon_transducer_inference.rst -.. py:module:: sofic.generators.epsilon_transducer_inference +.. transcssr.rst +.. py:module:: sofic.inference.cssr.transducer ********************* transCSSR Inference ********************* -transCSSR reconstructs an :doc:`ε-transducer ` from paired +transCSSR reconstructs an :doc:`ε-transducer <../generators/epsilon_transducer>` from paired input/output sample sequences, generalizing Causal-State Splitting Reconstruction :cite:`Shalizi2004` from a single process to an input/output channel :cite:`Barnett2015`. @@ -14,9 +14,9 @@ Causal states are equivalence classes of joint ``(input, output)`` pasts that induce the same conditional next-output law ``P(y | history, x)`` for every input symbol ``x``. Rare histories inherit their parent's state (controlled by ``min_count``); the split decision uses a G-test at significance ``alpha``, the -same test as :func:`~sofic.generators.epsilon_inference.morphs_differ` (with +same test as :func:`~sofic.inference.cssr.morphs_differ` (with Yates' continuity correction at one degree of freedom). -As in :func:`~sofic.generators.epsilon_inference.cssr`, ``test="exact"`` uses a +As in :func:`~sofic.inference.cssr.cssr`, ``test="exact"`` uses a Monte Carlo exact G-test for tables with small expected counts, and ``correction="bonferroni"`` divides ``alpha`` by the number of (history, input symbol) tests. ``Lmax="auto"`` sets the depth from the Markov diff --git a/sofic/automata/__init__.py b/sofic/automata/__init__.py index 4ed99ff..7587d5a 100644 --- a/sofic/automata/__init__.py +++ b/sofic/automata/__init__.py @@ -1,26 +1,6 @@ """Finite automata and transducers.""" # ``atomaton`` is an intentional pun on atomic automaton. -from sofic.automata.active import ( - AutomatonEquivalenceOracle, - EquivalenceOracle, - ExhaustiveEquivalenceOracle, - FunctionMealyOracle, - FunctionMembershipOracle, - LanguageMembershipOracle, - MealyEquivalenceOracle, - MealyExhaustiveEquivalenceOracle, - MealyMembershipOracle, - MembershipOracle, - RandomWalkEquivalenceOracle, - TransducerOutputOracle, - learn_dfa_from_language, - learn_dfa_lstar, - learn_dfa_ttt, - learn_mealy_from_transducer, - learn_mealy_lstar, -) -from sofic.automata.alergia import learn_pfa_alergia from sofic.automata.algorithms import ( MinimizationAlgorithm, complete, @@ -29,13 +9,12 @@ minimize, trim, ) -from sofic.automata.atomaton import Atomaton, AtomicAutomaton, MaximizedPrimeAtomaton from sofic.automata.base import LabeledAutomaton from sofic.automata.buchi import BuchiAutomaton +from sofic.automata.canonical.atomaton import Atomaton, AtomicAutomaton, MaximizedPrimeAtomaton +from sofic.automata.canonical.rfsa import CanonicalRFSA, ResidualFiniteStateAutomaton from sofic.automata.dfa import DFA -from sofic.automata.dfasat import learn_dfa_sat -from sofic.automata.edsm import learn_dfa_edsm -from sofic.automata.icdfa import ( +from sofic.automata.enumeration.icdfa import ( ICDFAString, count_flag_sequences, count_icdfa, @@ -52,7 +31,7 @@ string_from_flags, validate_icdfa_empty_string, ) -from sofic.automata.idfa import ( +from sofic.automata.enumeration.idfa import ( MISSING_TRANSITION, count_accessible_idfa, first_idfa_string, @@ -63,11 +42,31 @@ validate_idfa_string, ) from sofic.automata.languages import AutomatonLanguage, RegularLanguage -from sofic.automata.learning import learn_prime_atomaton_nlstar, learn_rfsa_from_language, learn_rfsa_nlstar -from sofic.automata.nfa import NFA -from sofic.automata.nwa import NestedWord, NestedWordAutomaton -from sofic.automata.observation import ObservationTable -from sofic.automata.papni import ( +from sofic.automata.learning.active import ( + AutomatonEquivalenceOracle, + EquivalenceOracle, + ExhaustiveEquivalenceOracle, + FunctionMealyOracle, + FunctionMembershipOracle, + LanguageMembershipOracle, + MealyEquivalenceOracle, + MealyExhaustiveEquivalenceOracle, + MealyMembershipOracle, + MembershipOracle, + RandomWalkEquivalenceOracle, + TransducerOutputOracle, + learn_dfa_from_language, + learn_dfa_lstar, + learn_dfa_ttt, + learn_mealy_from_transducer, + learn_mealy_lstar, +) +from sofic.automata.learning.alergia import learn_pfa_alergia +from sofic.automata.learning.dfasat import learn_dfa_sat +from sofic.automata.learning.edsm import learn_dfa_edsm +from sofic.automata.learning.nlstar import learn_prime_atomaton_nlstar, learn_rfsa_from_language, learn_rfsa_nlstar +from sofic.automata.learning.observation import ObservationTable +from sofic.automata.learning.papni import ( DyckAlphabet, is_well_matched, learn_sofic_dyck_shift_papni, @@ -75,9 +74,10 @@ papni_encode_samples, sofic_dyck_shift_from_papni_dfa, ) +from sofic.automata.learning.rpni import learn_dfa_rpni +from sofic.automata.nfa import NFA +from sofic.automata.nwa import NestedWord, NestedWordAutomaton from sofic.automata.regex import automaton_to_regex -from sofic.automata.rfsa import CanonicalRFSA, ResidualFiniteStateAutomaton -from sofic.automata.rpni import learn_dfa_rpni from sofic.automata.subsequential import ( SubsequentialTransducer, WeightedFiniteStateTransducer, diff --git a/sofic/automata/base.py b/sofic/automata/base.py index d75e0a9..6c3a226 100644 --- a/sofic/automata/base.py +++ b/sofic/automata/base.py @@ -47,7 +47,7 @@ def recognizes(self, word: Sequence[Any]) -> bool: def words_of_length(self, length: int) -> Iterator[tuple[Any, ...]]: """Yield accepted words of exactly ``length`` symbols.""" - from sofic.automata.enumeration import words_of_length + from sofic.automata.enumeration.words import words_of_length yield from words_of_length(self, length) @@ -57,7 +57,7 @@ def iter_language(self, max_length: int | None = None) -> Iterator[tuple[Any, .. If ``max_length`` is omitted, the iterator is unbounded and may not terminate for finite languages after yielding their last word. """ - from sofic.automata.enumeration import iter_language + from sofic.automata.enumeration.words import iter_language yield from iter_language(self, max_length=max_length) diff --git a/sofic/automata/canonical/__init__.py b/sofic/automata/canonical/__init__.py new file mode 100644 index 0000000..7b3bb5a --- /dev/null +++ b/sofic/automata/canonical/__init__.py @@ -0,0 +1,41 @@ +"""Canonical residual automata, atomata, and their duality.""" + +from sofic.automata.canonical.atomaton import ( + Atomaton, + AtomicAutomaton, + MaximizedPrimeAtomaton, + atomic_states, + is_atomic, +) +from sofic.automata.canonical.dual import dual_atomaton_from_rfsa, dual_rfsa_from_atomaton +from sofic.automata.canonical.residual import ( + ResidualTable, + atomaton_from_language, + canonical_rfsa_from_language, + maximized_prime_atomaton_from_language, + observation_to_atomaton, + observation_to_canonical_rfsa, + observation_to_maximized_prime_atomaton, + observation_to_minimal_dfa, +) +from sofic.automata.canonical.rfsa import CanonicalRFSA, ResidualFiniteStateAutomaton + +__all__ = [ + "Atomaton", + "AtomicAutomaton", + "CanonicalRFSA", + "MaximizedPrimeAtomaton", + "ResidualFiniteStateAutomaton", + "ResidualTable", + "atomaton_from_language", + "atomic_states", + "canonical_rfsa_from_language", + "dual_atomaton_from_rfsa", + "dual_rfsa_from_atomaton", + "is_atomic", + "maximized_prime_atomaton_from_language", + "observation_to_atomaton", + "observation_to_canonical_rfsa", + "observation_to_maximized_prime_atomaton", + "observation_to_minimal_dfa", +] diff --git a/sofic/automata/atomaton.py b/sofic/automata/canonical/atomaton.py similarity index 88% rename from sofic/automata/atomaton.py rename to sofic/automata/canonical/atomaton.py index 10e2b89..7e22abe 100644 --- a/sofic/automata/atomaton.py +++ b/sofic/automata/canonical/atomaton.py @@ -11,8 +11,8 @@ from sofic.exceptions import SoficValidationError if TYPE_CHECKING: - from sofic.automata.observation import ObservationTable - from sofic.automata.rfsa import CanonicalRFSA + from sofic.automata.canonical.rfsa import CanonicalRFSA + from sofic.automata.learning.observation import ObservationTable def atomic_states(nfa: NFA, *, alphabet: frozenset[Any] | None = None) -> frozenset[Hashable]: @@ -70,7 +70,7 @@ class Atomaton(AtomicAutomaton): @classmethod def from_language(cls, language: RegularLanguage | NFA, **kwargs: Any) -> Atomaton: - from sofic.automata.canonical_extraction import atomaton_from_language + from sofic.automata.canonical.residual import atomaton_from_language return atomaton_from_language(language) @@ -92,25 +92,25 @@ class MaximizedPrimeAtomaton(NFA): @classmethod def from_language(cls, language: RegularLanguage | NFA, **kwargs: Any) -> MaximizedPrimeAtomaton: - from sofic.automata.canonical_extraction import maximized_prime_atomaton_from_language + from sofic.automata.canonical.residual import maximized_prime_atomaton_from_language return maximized_prime_atomaton_from_language(language) @classmethod def from_observation_table(cls, table: ObservationTable, **kwargs: Any) -> MaximizedPrimeAtomaton: - from sofic.automata.canonical_extraction import observation_to_maximized_prime_atomaton + from sofic.automata.canonical.residual import observation_to_maximized_prime_atomaton return observation_to_maximized_prime_atomaton(table) @classmethod def from_canonical_rfsa(cls, rfsa: CanonicalRFSA, **kwargs: Any) -> MaximizedPrimeAtomaton: """Return the maximized prime átomaton of the language ``rfsa`` recognizes.""" - from sofic.automata.canonical_extraction import maximized_prime_atomaton_from_language + from sofic.automata.canonical.residual import maximized_prime_atomaton_from_language return maximized_prime_atomaton_from_language(rfsa) def dual(self) -> CanonicalRFSA: """Return the reverse automaton: the canonical RFSA of the reversed language.""" - from sofic.automata.canonical_dual import dual_rfsa_from_atomaton + from sofic.automata.canonical.dual import dual_rfsa_from_atomaton return dual_rfsa_from_atomaton(self) diff --git a/sofic/automata/canonical_dual.py b/sofic/automata/canonical/dual.py similarity index 76% rename from sofic/automata/canonical_dual.py rename to sofic/automata/canonical/dual.py index 4ba8469..4a4aaf4 100644 --- a/sofic/automata/canonical_dual.py +++ b/sofic/automata/canonical/dual.py @@ -7,19 +7,19 @@ from __future__ import annotations -from sofic.automata.atomaton import MaximizedPrimeAtomaton -from sofic.automata.rfsa import CanonicalRFSA +from sofic.automata.canonical.atomaton import MaximizedPrimeAtomaton +from sofic.automata.canonical.rfsa import CanonicalRFSA def dual_atomaton_from_rfsa(rfsa: CanonicalRFSA) -> MaximizedPrimeAtomaton: """Reverse the canonical RFSA of ``L`` into the maximized prime átomaton of ``L^R``.""" - from sofic.automata.canonical_extraction import _reverse_into + from sofic.automata.canonical.residual import _reverse_into return _reverse_into(MaximizedPrimeAtomaton, rfsa) def dual_rfsa_from_atomaton(atomaton: MaximizedPrimeAtomaton) -> CanonicalRFSA: """Reverse the maximized prime átomaton of ``L`` into the canonical RFSA of ``L^R``.""" - from sofic.automata.canonical_extraction import _reverse_into + from sofic.automata.canonical.residual import _reverse_into return _reverse_into(CanonicalRFSA, atomaton) diff --git a/sofic/automata/canonical_extraction.py b/sofic/automata/canonical/residual.py similarity index 97% rename from sofic/automata/canonical_extraction.py rename to sofic/automata/canonical/residual.py index 61a3d90..b463e38 100644 --- a/sofic/automata/canonical_extraction.py +++ b/sofic/automata/canonical/residual.py @@ -20,13 +20,13 @@ from dataclasses import dataclass from typing import Any -from sofic.automata.atomaton import Atomaton, MaximizedPrimeAtomaton from sofic.automata.base import LabeledAutomaton +from sofic.automata.canonical.atomaton import Atomaton, MaximizedPrimeAtomaton +from sofic.automata.canonical.rfsa import CanonicalRFSA from sofic.automata.dfa import DFA from sofic.automata.languages.base import AutomatonLanguage, RegularLanguage, as_language +from sofic.automata.learning.observation import ObservationTable from sofic.automata.nfa import NFA -from sofic.automata.observation import ObservationTable -from sofic.automata.rfsa import CanonicalRFSA from sofic.graph import ATTR_SYMBOL from sofic.states import sequential_labels @@ -283,7 +283,7 @@ def observation_to_canonical_rfsa(table: ObservationTable) -> CanonicalRFSA: States are the rows of the access words that are prime among all rows of the table, exactly as in the NL\* hypothesis :cite:`Bollig2009`. """ - from sofic.automata.learning import _hypothesis, _primes + from sofic.automata.learning.nlstar import _hypothesis, _primes experiments = sorted(table.experiments, key=lambda w: (len(w), w)) if () in experiments: diff --git a/sofic/automata/rfsa.py b/sofic/automata/canonical/rfsa.py similarity index 80% rename from sofic/automata/rfsa.py rename to sofic/automata/canonical/rfsa.py index 846b867..1a0bf48 100644 --- a/sofic/automata/rfsa.py +++ b/sofic/automata/canonical/rfsa.py @@ -8,8 +8,8 @@ from sofic.automata.nfa import NFA if TYPE_CHECKING: - from sofic.automata.atomaton import MaximizedPrimeAtomaton - from sofic.automata.observation import ObservationTable + from sofic.automata.canonical.atomaton import MaximizedPrimeAtomaton + from sofic.automata.learning.observation import ObservationTable class ResidualFiniteStateAutomaton(NFA): @@ -18,7 +18,7 @@ class ResidualFiniteStateAutomaton(NFA): def validate(self) -> None: super().validate() from sofic.automata.algorithms import equivalent - from sofic.automata.canonical_extraction import ResidualTable + from sofic.automata.canonical.residual import ResidualTable if not self.initial_states: return @@ -40,18 +40,18 @@ class CanonicalRFSA(ResidualFiniteStateAutomaton): @classmethod def from_language(cls, language: RegularLanguage | NFA, **kwargs: Any) -> CanonicalRFSA: - from sofic.automata.canonical_extraction import canonical_rfsa_from_language + from sofic.automata.canonical.residual import canonical_rfsa_from_language return canonical_rfsa_from_language(language) @classmethod def from_observation_table(cls, table: ObservationTable, **kwargs: Any) -> CanonicalRFSA: - from sofic.automata.canonical_extraction import observation_to_canonical_rfsa + from sofic.automata.canonical.residual import observation_to_canonical_rfsa return observation_to_canonical_rfsa(table) def dual(self) -> MaximizedPrimeAtomaton: """Return the reverse automaton: the maximized prime átomaton of the reversed language.""" - from sofic.automata.canonical_dual import dual_atomaton_from_rfsa + from sofic.automata.canonical.dual import dual_atomaton_from_rfsa return dual_atomaton_from_rfsa(self) diff --git a/sofic/automata/enumeration/__init__.py b/sofic/automata/enumeration/__init__.py new file mode 100644 index 0000000..e09ea83 --- /dev/null +++ b/sofic/automata/enumeration/__init__.py @@ -0,0 +1,58 @@ +"""Enumeration of words and of initially connected / accessible DFAs.""" + +from sofic.automata.enumeration.icdfa import ( + ICDFAString, + count_flag_sequences, + count_icdfa, + count_icdfa_empty, + dfa_to_icdfa_string, + first_icdfa_empty_string, + flags_from_string, + icdfa_string_to_dfa, + iter_icdfa, + iter_icdfa_empty_strings, + last_icdfa_empty_string, + next_flags, + next_icdfa_empty_string, + string_from_flags, + validate_icdfa_empty_string, +) +from sofic.automata.enumeration.idfa import ( + MISSING_TRANSITION, + count_accessible_idfa, + first_idfa_string, + iter_idfa_strings, + rank_idfa_string, + reroot_idfa_string, + unrank_idfa_string, + validate_idfa_string, +) +from sofic.automata.enumeration.words import iter_language, words_of_length + +__all__ = [ + "ICDFAString", + "MISSING_TRANSITION", + "count_accessible_idfa", + "count_flag_sequences", + "count_icdfa", + "count_icdfa_empty", + "dfa_to_icdfa_string", + "first_icdfa_empty_string", + "first_idfa_string", + "flags_from_string", + "icdfa_string_to_dfa", + "iter_icdfa", + "iter_icdfa_empty_strings", + "iter_idfa_strings", + "iter_language", + "last_icdfa_empty_string", + "next_flags", + "next_icdfa_empty_string", + "rank_idfa_string", + "reroot_idfa_string", + "string_from_flags", + "unrank_idfa_string", + "validate_icdfa_empty_string", + "validate_idfa_string", + "words_of_length", +] diff --git a/sofic/automata/icdfa.py b/sofic/automata/enumeration/icdfa.py similarity index 100% rename from sofic/automata/icdfa.py rename to sofic/automata/enumeration/icdfa.py diff --git a/sofic/automata/idfa.py b/sofic/automata/enumeration/idfa.py similarity index 99% rename from sofic/automata/idfa.py rename to sofic/automata/enumeration/idfa.py index ad7db6d..56adfd7 100644 --- a/sofic/automata/idfa.py +++ b/sofic/automata/enumeration/idfa.py @@ -11,7 +11,7 @@ from collections.abc import Iterator, Sequence from functools import cache -from sofic.automata.icdfa import ( +from sofic.automata.enumeration.icdfa import ( _upper_bound_at, _validate_flags, next_flags, diff --git a/sofic/automata/enumeration.py b/sofic/automata/enumeration/words.py similarity index 100% rename from sofic/automata/enumeration.py rename to sofic/automata/enumeration/words.py diff --git a/sofic/automata/languages/atoms.py b/sofic/automata/languages/atoms.py index 38de8a9..fa625fe 100644 --- a/sofic/automata/languages/atoms.py +++ b/sofic/automata/languages/atoms.py @@ -15,7 +15,7 @@ def _reversed_table(language: RegularLanguage): - from sofic.automata.canonical_extraction import ResidualTable + from sofic.automata.canonical.residual import ResidualTable lang = as_language(language) # type: ignore[arg-type] if not isinstance(lang, AutomatonLanguage): diff --git a/sofic/automata/languages/residuals.py b/sofic/automata/languages/residuals.py index 5e26d79..c412f08 100644 --- a/sofic/automata/languages/residuals.py +++ b/sofic/automata/languages/residuals.py @@ -19,7 +19,7 @@ def prime_residuals(language: RegularLanguage) -> frozenset[RegularLanguage]: """Return the prime residuals of ``language``.""" lang = as_language(language) # type: ignore[arg-type] if isinstance(lang, AutomatonLanguage): - from sofic.automata.canonical_extraction import ResidualTable + from sofic.automata.canonical.residual import ResidualTable table = ResidualTable.from_automaton(lang.automaton) return frozenset(AutomatonLanguage(table.residual_automaton(q)) for q in table.prime_states()) diff --git a/sofic/automata/learning/__init__.py b/sofic/automata/learning/__init__.py new file mode 100644 index 0000000..8a1c4cc --- /dev/null +++ b/sofic/automata/learning/__init__.py @@ -0,0 +1,69 @@ +"""Passive and active automaton learning algorithms.""" + +from sofic.automata.learning.active import ( + AutomatonEquivalenceOracle, + EquivalenceOracle, + ExhaustiveEquivalenceOracle, + FunctionMealyOracle, + FunctionMembershipOracle, + LanguageMembershipOracle, + MealyEquivalenceOracle, + MealyExhaustiveEquivalenceOracle, + MealyMembershipOracle, + MembershipOracle, + RandomWalkEquivalenceOracle, + TransducerOutputOracle, + learn_dfa_from_language, + learn_dfa_lstar, + learn_dfa_ttt, + learn_mealy_from_transducer, + learn_mealy_lstar, +) +from sofic.automata.learning.alergia import learn_pfa_alergia +from sofic.automata.learning.dfasat import learn_dfa_sat +from sofic.automata.learning.edsm import learn_dfa_edsm +from sofic.automata.learning.nlstar import learn_prime_atomaton_nlstar, learn_rfsa_from_language, learn_rfsa_nlstar +from sofic.automata.learning.observation import ObservationTable +from sofic.automata.learning.papni import ( + DyckAlphabet, + is_well_matched, + learn_sofic_dyck_shift_papni, + papni_encode, + papni_encode_samples, + sofic_dyck_shift_from_papni_dfa, +) +from sofic.automata.learning.rpni import learn_dfa_rpni + +__all__ = [ + "AutomatonEquivalenceOracle", + "DyckAlphabet", + "EquivalenceOracle", + "ExhaustiveEquivalenceOracle", + "FunctionMealyOracle", + "FunctionMembershipOracle", + "LanguageMembershipOracle", + "MealyEquivalenceOracle", + "MealyExhaustiveEquivalenceOracle", + "MealyMembershipOracle", + "MembershipOracle", + "ObservationTable", + "RandomWalkEquivalenceOracle", + "TransducerOutputOracle", + "is_well_matched", + "learn_dfa_edsm", + "learn_dfa_from_language", + "learn_dfa_lstar", + "learn_dfa_rpni", + "learn_dfa_sat", + "learn_dfa_ttt", + "learn_mealy_from_transducer", + "learn_mealy_lstar", + "learn_pfa_alergia", + "learn_prime_atomaton_nlstar", + "learn_rfsa_from_language", + "learn_rfsa_nlstar", + "learn_sofic_dyck_shift_papni", + "papni_encode", + "papni_encode_samples", + "sofic_dyck_shift_from_papni_dfa", +] diff --git a/sofic/automata/active.py b/sofic/automata/learning/active.py similarity index 99% rename from sofic/automata/active.py rename to sofic/automata/learning/active.py index 55fc281..386dd71 100644 --- a/sofic/automata/active.py +++ b/sofic/automata/learning/active.py @@ -14,7 +14,7 @@ :cite:`KearnsVazirani1994,Isberner2014`. These complement the NL\* canonical-RFSA learner -(:func:`sofic.automata.learning.learn_rfsa_nlstar`). +(:func:`sofic.automata.learning.nlstar.learn_rfsa_nlstar`). """ from __future__ import annotations diff --git a/sofic/automata/alergia.py b/sofic/automata/learning/alergia.py similarity index 99% rename from sofic/automata/alergia.py rename to sofic/automata/learning/alergia.py index 1df3f26..2adc9a0 100644 --- a/sofic/automata/alergia.py +++ b/sofic/automata/learning/alergia.py @@ -5,7 +5,7 @@ whenever a Hoeffding-bound test cannot distinguish their outgoing (and recursive) transition statistics. It is the stochastic, unlabeled counterpart of RPNI/EDSM and a state-merging alternative to Causal-State Splitting Reconstruction -(:func:`sofic.generators.epsilon_inference.cssr`). +(:func:`sofic.inference.cssr.process.cssr`). The learned automaton is returned as a :class:`~sofic.generators.pfa.ProbabilisticFiniteAutomaton` describing the diff --git a/sofic/automata/dfasat.py b/sofic/automata/learning/dfasat.py similarity index 95% rename from sofic/automata/dfasat.py rename to sofic/automata/learning/dfasat.py index a3cdac5..cbe6858 100644 --- a/sofic/automata/dfasat.py +++ b/sofic/automata/learning/dfasat.py @@ -7,7 +7,7 @@ into a graph-colouring SAT instance -- one Boolean per (APTA node, colour), plus accepting-colour and transition variables -- and an external SAT solver is asked whether a ``k``-colouring exists. The search runs ``k`` upward from a lower bound -to the EDSM upper bound (:func:`sofic.automata.edsm.learn_dfa_edsm`); the first +to the EDSM upper bound (:func:`sofic.automata.learning.edsm.learn_dfa_edsm`); the first satisfiable ``k`` is the provably minimal DFA. This requires the optional `python-sat `_ dependency @@ -20,7 +20,7 @@ from typing import Any from sofic.automata.dfa import DFA -from sofic.automata.edsm import _ACCEPT, _REJECT, _build_apta, _collect_alphabet, learn_dfa_edsm +from sofic.automata.learning.edsm import _ACCEPT, _REJECT, _build_apta, _collect_alphabet, learn_dfa_edsm __all__ = ["learn_dfa_sat"] @@ -156,7 +156,7 @@ def learn_dfa_sat( Smallest state count to try (default ``1``). upper_bound Largest state count to try. Defaults to the number of states of the - EDSM hypothesis (:func:`sofic.automata.edsm.learn_dfa_edsm`), which is a + EDSM hypothesis (:func:`sofic.automata.learning.edsm.learn_dfa_edsm`), which is a valid upper bound on the minimum. solver_name Any `python-sat `_ solver name (default diff --git a/sofic/automata/edsm.py b/sofic/automata/learning/edsm.py similarity index 99% rename from sofic/automata/edsm.py rename to sofic/automata/learning/edsm.py index e2b8e32..8c46c5f 100644 --- a/sofic/automata/edsm.py +++ b/sofic/automata/learning/edsm.py @@ -1,6 +1,6 @@ """Passive DFA learning via Evidence-Driven State Merging (blue-fringe). -EDSM upgrades the greedy RPNI merge order (:func:`sofic.automata.rpni.learn_dfa_rpni`) +EDSM upgrades the greedy RPNI merge order (:func:`sofic.automata.learning.rpni.learn_dfa_rpni`) with the evidence-driven, red/blue "blue-fringe" strategy that won the Abbadingo One competition :cite:`Lang1998`. Starting from the augmented prefix-tree acceptor of the labeled sample, it maintains a set of confirmed *red* states and diff --git a/sofic/automata/learning.py b/sofic/automata/learning/nlstar.py similarity index 96% rename from sofic/automata/learning.py rename to sofic/automata/learning/nlstar.py index d461d52..a73114a 100644 --- a/sofic/automata/learning.py +++ b/sofic/automata/learning/nlstar.py @@ -23,7 +23,10 @@ from collections.abc import Iterable, Sequence from typing import Any -from sofic.automata.active import ( +from sofic.automata.base import LabeledAutomaton +from sofic.automata.canonical.atomaton import MaximizedPrimeAtomaton +from sofic.automata.canonical.rfsa import CanonicalRFSA +from sofic.automata.learning.active import ( AutomatonEquivalenceOracle, EquivalenceOracle, ExhaustiveEquivalenceOracle, @@ -31,9 +34,6 @@ MembershipOracle, _MembershipCache, ) -from sofic.automata.atomaton import MaximizedPrimeAtomaton -from sofic.automata.base import LabeledAutomaton -from sofic.automata.rfsa import CanonicalRFSA Word = tuple[Any, ...] Row = tuple[bool, ...] @@ -174,7 +174,7 @@ def learn_prime_atomaton_nlstar( :math:`L^R` :cite:`MaarandTamm2022`, so NL\* is run against reversed membership and equivalence oracles and its result reversed. """ - from sofic.automata.canonical_dual import dual_atomaton_from_rfsa + from sofic.automata.canonical.dual import dual_atomaton_from_rfsa reversed_rfsa = learn_rfsa_nlstar( alphabet, _ReversedMembership(membership), _ReversedEquivalence(equivalence), max_rounds=max_rounds diff --git a/sofic/automata/observation.py b/sofic/automata/learning/observation.py similarity index 68% rename from sofic/automata/observation.py rename to sofic/automata/learning/observation.py index 3bfd361..a0b6637 100644 --- a/sofic/automata/observation.py +++ b/sofic/automata/learning/observation.py @@ -6,9 +6,9 @@ from typing import TYPE_CHECKING, Any if TYPE_CHECKING: - from sofic.automata.atomaton import Atomaton, MaximizedPrimeAtomaton + from sofic.automata.canonical.atomaton import Atomaton, MaximizedPrimeAtomaton + from sofic.automata.canonical.rfsa import CanonicalRFSA from sofic.automata.dfa import DFA - from sofic.automata.rfsa import CanonicalRFSA @dataclass @@ -20,21 +20,21 @@ class ObservationTable: membership: dict[tuple[Any, ...], bool] = field(default_factory=dict) def to_minimal_dfa(self) -> DFA: - from sofic.automata.canonical_extraction import observation_to_minimal_dfa + from sofic.automata.canonical.residual import observation_to_minimal_dfa return observation_to_minimal_dfa(self) def to_canonical_rfsa(self) -> CanonicalRFSA: - from sofic.automata.canonical_extraction import observation_to_canonical_rfsa + from sofic.automata.canonical.residual import observation_to_canonical_rfsa return observation_to_canonical_rfsa(self) def to_atomaton(self) -> Atomaton: - from sofic.automata.canonical_extraction import observation_to_atomaton + from sofic.automata.canonical.residual import observation_to_atomaton return observation_to_atomaton(self) def to_maximized_prime_atomaton(self) -> MaximizedPrimeAtomaton: - from sofic.automata.canonical_extraction import observation_to_maximized_prime_atomaton + from sofic.automata.canonical.residual import observation_to_maximized_prime_atomaton return observation_to_maximized_prime_atomaton(self) diff --git a/sofic/automata/papni.py b/sofic/automata/learning/papni.py similarity index 99% rename from sofic/automata/papni.py rename to sofic/automata/learning/papni.py index d1c367f..163a8af 100644 --- a/sofic/automata/papni.py +++ b/sofic/automata/learning/papni.py @@ -7,7 +7,7 @@ from typing import Any from sofic.automata.dfa import DFA -from sofic.automata.rpni import learn_dfa_rpni +from sofic.automata.learning.rpni import learn_dfa_rpni from sofic.graph import ATTR_KIND, ATTR_SYMBOL, KIND_CALL, KIND_INTERNAL, KIND_RETURN from sofic.shifts.sofic_dyck import SoficDyckShift, TransitionRef, transition_ref diff --git a/sofic/automata/rpni.py b/sofic/automata/learning/rpni.py similarity index 100% rename from sofic/automata/rpni.py rename to sofic/automata/learning/rpni.py diff --git a/sofic/automata/vpa.py b/sofic/automata/vpa.py deleted file mode 100644 index c70c828..0000000 --- a/sofic/automata/vpa.py +++ /dev/null @@ -1,1331 +0,0 @@ -"""Visibly pushdown automata.""" - -from __future__ import annotations - -from collections import deque -from collections.abc import Hashable, Iterable, Mapping, Sequence -from typing import Any - -from sofic.automata import vpa_constructions as vc -from sofic.base import StateMachine -from sofic.exceptions import NonDeterministicError, NonWellMatchedLanguageError -from sofic.graph import ( - ATTR_KIND, - ATTR_STACK_SYMBOL, - ATTR_SYMBOL, - KIND_CALL, - KIND_INTERNAL, - KIND_RETURN, -) - -_MISSING = object() - - -class VisiblyPushdownAutomaton(StateMachine): - """Standard 1-stack VPA with call / return / internal input partition.""" - - input_alphabet: frozenset[Any] - call_alphabet: frozenset[Any] - return_alphabet: frozenset[Any] - internal_alphabet: frozenset[Any] - stack_alphabet: frozenset[Any] - bottom_stack_symbol: Any | None - initial_state: Hashable | None - accepting_states: frozenset[Hashable] - - def __init__( - self, - input_alphabet: frozenset[Any] | None = None, - call_alphabet: frozenset[Any] | None = None, - return_alphabet: frozenset[Any] | None = None, - internal_alphabet: frozenset[Any] | None = None, - stack_alphabet: frozenset[Any] | None = None, - bottom_stack_symbol: Any | None = None, - initial_state: Hashable | None = None, - accepting_states: frozenset[Hashable] | None = None, - **kwargs: Any, - ) -> None: - super().__init__(**kwargs) - self.call_alphabet = call_alphabet if call_alphabet is not None else frozenset() - self.return_alphabet = return_alphabet if return_alphabet is not None else frozenset() - self.internal_alphabet = internal_alphabet if internal_alphabet is not None else frozenset() - self.stack_alphabet = stack_alphabet if stack_alphabet is not None else frozenset() - self.bottom_stack_symbol = bottom_stack_symbol - self.input_alphabet = ( - input_alphabet - if input_alphabet is not None - else self.call_alphabet | self.return_alphabet | self.internal_alphabet - ) - self.initial_state = initial_state - self.accepting_states = accepting_states if accepting_states is not None else frozenset() - - def validate(self) -> None: - partition = self.call_alphabet | self.return_alphabet | self.internal_alphabet - self._require( - len(self.call_alphabet) + len(self.return_alphabet) + len(self.internal_alphabet) == len(partition), - "call, return, and internal alphabets must be disjoint", - ) - self._require(partition == self.input_alphabet, "input alphabet must equal partition of call/return/internal") - if self.bottom_stack_symbol is not None: - self._require( - self.bottom_stack_symbol in self.stack_alphabet, "bottom_stack_symbol must be in stack alphabet" - ) - if self.initial_state is not None: - self._require(self.graph.has_state(self.initial_state), "missing initial state") - for state in self.accepting_states: - self._require(self.graph.has_state(state), f"missing accepting state {state!r}") - for transition in self.transitions(): - kind = transition.data.get(ATTR_KIND) - symbol = transition.data.get(ATTR_SYMBOL) - self._require(kind in {KIND_CALL, KIND_RETURN, KIND_INTERNAL}, f"invalid VPA kind {kind!r}") - if symbol is not None: - if kind == KIND_CALL: - self._require(symbol in self.call_alphabet, f"{symbol!r} not in call alphabet") - stack_sym = transition.data.get(ATTR_STACK_SYMBOL) - self._require(stack_sym in self.stack_alphabet, "call edge requires stack_symbol in stack alphabet") - self._require( - stack_sym != self.bottom_stack_symbol, - "call edge cannot push the bottom_stack_symbol", - ) - elif kind == KIND_RETURN: - self._require(symbol in self.return_alphabet, f"{symbol!r} not in return alphabet") - stack_sym = transition.data.get(ATTR_STACK_SYMBOL) - if stack_sym is not None: - self._require(stack_sym in self.stack_alphabet, "return stack_symbol must be in stack alphabet") - else: - self._require(symbol in self.internal_alphabet, f"{symbol!r} not in internal alphabet") - - def add_call_transition( - self, - source: Hashable, - target: Hashable, - symbol: Any, - stack_symbol: Any, - **attrs: Any, - ) -> int: - """Add a call transition that pushes ``stack_symbol``.""" - data = {**attrs, ATTR_KIND: KIND_CALL, ATTR_SYMBOL: symbol, ATTR_STACK_SYMBOL: stack_symbol} - return self.graph.add_transition(source, target, **data) - - def add_return_transition( - self, - source: Hashable, - target: Hashable, - symbol: Any, - stack_symbol: Any | None = None, - **attrs: Any, - ) -> int: - """Add a return transition. - - If ``stack_symbol`` is omitted, the transition is a wildcard: it fires - on every stack symbol, and also on the empty stack (leaving it empty) - when the VPA has a ``bottom_stack_symbol``. A return guarded by the - bottom symbol fires only on the empty stack. - """ - data = {**attrs, ATTR_KIND: KIND_RETURN, ATTR_SYMBOL: symbol} - if stack_symbol is not None: - data[ATTR_STACK_SYMBOL] = stack_symbol - return self.graph.add_transition(source, target, **data) - - def add_internal_transition(self, source: Hashable, target: Hashable, symbol: Any, **attrs: Any) -> int: - """Add an internal transition.""" - data = {**attrs, ATTR_KIND: KIND_INTERNAL, ATTR_SYMBOL: symbol} - return self.graph.add_transition(source, target, **data) - - def call_transition_map(self) -> dict[tuple[Hashable, Any], tuple[Hashable, Any]]: - """Return deterministic call transitions keyed by ``(state, symbol)``.""" - result: dict[tuple[Hashable, Any], tuple[Hashable, Any]] = {} - for transition in self.transitions(): - if transition.data.get(ATTR_KIND) != KIND_CALL: - continue - symbol = transition.data.get(ATTR_SYMBOL) - if symbol is None: - continue - key = (transition.source, symbol) - value = (transition.target, transition.data.get(ATTR_STACK_SYMBOL)) - if key in result and result[key] != value: - raise NonDeterministicError(f"non-deterministic call transition on {key}") - result[key] = value - return result - - def return_transition_map(self) -> dict[tuple[Hashable, Any, Any | None], Hashable]: - """Return deterministic return transitions keyed by ``(state, symbol, stack_symbol)``.""" - result: dict[tuple[Hashable, Any, Any | None], Hashable] = {} - for transition in self.transitions(): - if transition.data.get(ATTR_KIND) != KIND_RETURN: - continue - symbol = transition.data.get(ATTR_SYMBOL) - if symbol is None: - continue - key = (transition.source, symbol, transition.data.get(ATTR_STACK_SYMBOL)) - value = transition.target - if key in result and result[key] != value: - raise NonDeterministicError(f"non-deterministic return transition on {key}") - result[key] = value - return result - - def internal_transition_map(self) -> dict[tuple[Hashable, Any], Hashable]: - """Return deterministic internal transitions keyed by ``(state, symbol)``.""" - result: dict[tuple[Hashable, Any], Hashable] = {} - for transition in self.transitions(): - if transition.data.get(ATTR_KIND) != KIND_INTERNAL: - continue - symbol = transition.data.get(ATTR_SYMBOL) - if symbol is None: - continue - key = (transition.source, symbol) - value = transition.target - if key in result and result[key] != value: - raise NonDeterministicError(f"non-deterministic internal transition on {key}") - result[key] = value - return result - - def recognizes(self, word: Sequence[Any]) -> bool: - from sofic.automata.vpa_simulation import recognizes_vpa - - return recognizes_vpa(self, word) - - def union(self, other: VisiblyPushdownAutomaton) -> VisiblyPushdownAutomaton: - """Return a VPA recognizing the union with ``other``.""" - return _binary(vc.union, self, other) - - def intersection(self, other: VisiblyPushdownAutomaton) -> VisiblyPushdownAutomaton: - """Return a VPA recognizing the intersection with ``other`` (synchronized product).""" - return _binary(vc.intersection, self, other) - - def complement(self) -> DeterministicVisiblyPushdownAutomaton: - """Return a deterministic VPA recognizing the complement over this visible alphabet.""" - return vc.denormalize(vc.complement(vc.normalize(self)), DeterministicVisiblyPushdownAutomaton) - - def difference(self, other: VisiblyPushdownAutomaton) -> VisiblyPushdownAutomaton: - """Return a VPA recognizing this language minus ``other``.""" - return _binary(vc.difference, self, other) - - def concat(self, other: VisiblyPushdownAutomaton) -> VisiblyPushdownAutomaton: - """Return a VPA recognizing concatenation with ``other``. - - Each factor is read from an empty stack of its own: a return in the right - factor that would pop a pending call of the left factor is a pending - return of the right factor. - """ - return _binary(vc.concat, self, other) - - def kleene_star(self) -> VisiblyPushdownAutomaton: - """Return a VPA recognizing the Kleene star; each factor starts from its own empty stack.""" - return vc.denormalize(vc.kleene_star(vc.normalize(self))) - - def determinize(self) -> DeterministicVisiblyPushdownAutomaton: - """Return an equivalent complete deterministic VPA :cite:`AlurMadhusudan2009`.""" - return vc.determinize_vpa(self) - - def is_empty(self) -> bool: - """Return whether the language is empty.""" - return vc.is_empty(vc.normalize(self)) - - def accepted_word(self) -> tuple[Any, ...] | None: - """Return a short accepted word, or ``None`` when the language is empty.""" - return vc.accepted_word(vc.normalize(self)) - - def is_universal(self) -> bool: - """Return whether every word over the visible alphabet is accepted.""" - return vc.is_empty(vc.complement(vc.normalize(self))) - - def includes(self, other: VisiblyPushdownAutomaton) -> bool: - """Return whether ``other``'s language is contained in this one.""" - return vc.is_empty(vc.difference(vc.normalize(other), vc.normalize(self))) - - def equivalent(self, other: VisiblyPushdownAutomaton) -> bool: - """Return whether both VPAs recognize the same language.""" - return self.includes(other) and other.includes(self) - - def has_unmatched_word(self) -> bool: - """Return whether some accepted word has a pending call or a pending return.""" - return vc.has_unmatched_word(vc.normalize(self)) - - -def _binary(operation, left: VisiblyPushdownAutomaton, right: VisiblyPushdownAutomaton) -> VisiblyPushdownAutomaton: - return vc.denormalize(operation(vc.normalize(left), vc.normalize(right))) - - -class DeterministicVisiblyPushdownAutomaton(VisiblyPushdownAutomaton): - """VPA with at most one enabled transition for each visible configuration.""" - - def validate(self) -> None: - super().validate() - if self.initial_state is None: - raise NonDeterministicError("deterministic VPA requires an initial state") - self._check_determinism() - - def _check_determinism(self) -> None: - call_keys: set[tuple[Hashable, Any]] = set() - internal_keys: set[tuple[Hashable, Any]] = set() - return_keys: set[tuple[Hashable, Any, Any]] = set() - wildcard_returns: set[tuple[Hashable, Any]] = set() - - for transition in self.transitions(): - kind = transition.data.get(ATTR_KIND) - symbol = transition.data.get(ATTR_SYMBOL) - if symbol is None: - continue - if kind == KIND_CALL: - key = (transition.source, symbol) - if key in call_keys: - raise NonDeterministicError(f"non-deterministic call transition on {key}") - call_keys.add(key) - elif kind == KIND_INTERNAL: - key = (transition.source, symbol) - if key in internal_keys: - raise NonDeterministicError(f"non-deterministic internal transition on {key}") - internal_keys.add(key) - elif kind == KIND_RETURN: - stack_symbol = transition.data.get(ATTR_STACK_SYMBOL) - wildcard_key = (transition.source, symbol) - if stack_symbol is None: - if wildcard_key in wildcard_returns: - raise NonDeterministicError(f"duplicate wildcard return transition on {wildcard_key}") - if any(source == transition.source and ret == symbol for source, ret, _stack in return_keys): - raise NonDeterministicError( - f"wildcard return overlaps guarded return on {wildcard_key}", - ) - wildcard_returns.add(wildcard_key) - else: - key = (transition.source, symbol, stack_symbol) - if wildcard_key in wildcard_returns: - raise NonDeterministicError( - f"guarded return overlaps wildcard return on {wildcard_key}", - ) - if key in return_keys: - raise NonDeterministicError(f"duplicate guarded return transition on {key}") - return_keys.add(key) - - def add_call_transition( - self, - source: Hashable, - target: Hashable, - symbol: Any, - stack_symbol: Any, - **attrs: Any, - ) -> int: - for transition in self.graph.out_transitions(source): - if transition.data.get(ATTR_KIND) == KIND_CALL and transition.data.get(ATTR_SYMBOL) == symbol: - raise NonDeterministicError(f"non-deterministic call transition on {(source, symbol)}") - return super().add_call_transition(source, target, symbol, stack_symbol, **attrs) - - def add_return_transition( - self, - source: Hashable, - target: Hashable, - symbol: Any, - stack_symbol: Any | None = None, - **attrs: Any, - ) -> int: - for transition in self.graph.out_transitions(source): - if transition.data.get(ATTR_KIND) != KIND_RETURN or transition.data.get(ATTR_SYMBOL) != symbol: - continue - existing_stack = transition.data.get(ATTR_STACK_SYMBOL) - if existing_stack is None or stack_symbol is None or existing_stack == stack_symbol: - raise NonDeterministicError(f"non-deterministic return transition on {(source, symbol)}") - return super().add_return_transition(source, target, symbol, stack_symbol, **attrs) - - def add_internal_transition(self, source: Hashable, target: Hashable, symbol: Any, **attrs: Any) -> int: - for transition in self.graph.out_transitions(source): - if transition.data.get(ATTR_KIND) == KIND_INTERNAL and transition.data.get(ATTR_SYMBOL) == symbol: - raise NonDeterministicError(f"non-deterministic internal transition on {(source, symbol)}") - return super().add_internal_transition(source, target, symbol, **attrs) - - def call_successor(self, state: Hashable, symbol: Any) -> tuple[Hashable, Any] | None: - """Return ``(target, pushed_stack_symbol)`` for a deterministic call.""" - return self.call_transition_map().get((state, symbol)) - - def internal_successor(self, state: Hashable, symbol: Any) -> Hashable | None: - """Return the deterministic internal successor, if present.""" - return self.internal_transition_map().get((state, symbol)) - - def return_successor(self, state: Hashable, symbol: Any, stack_symbol: Any) -> Hashable | None: - """Return the deterministic return successor for ``stack_symbol``, if present.""" - transitions = self.return_transition_map() - explicit = transitions.get((state, symbol, stack_symbol), _MISSING) - if explicit is not _MISSING: - return explicit - wildcard = transitions.get((state, symbol, None), _MISSING) - if wildcard is not _MISSING: - return wildcard - return None - - @classmethod - def from_vpa(cls, vpa: VisiblyPushdownAutomaton) -> DeterministicVisiblyPushdownAutomaton: - """Return ``vpa`` as a deterministic VPA, determinizing it when needed. - - An already deterministic ``vpa`` is copied with its states unchanged; - otherwise the summary construction of :cite:`AlurMadhusudan2009` is - applied (see :func:`~sofic.automata.vpa_constructions.determinize`). - """ - if vpa.initial_state is None or not _is_deterministic(vpa): - return vc.determinize_vpa(vpa) - result = cls( - input_alphabet=vpa.input_alphabet, - call_alphabet=vpa.call_alphabet, - return_alphabet=vpa.return_alphabet, - internal_alphabet=vpa.internal_alphabet, - stack_alphabet=vpa.stack_alphabet, - bottom_stack_symbol=vpa.bottom_stack_symbol, - initial_state=vpa.initial_state, - accepting_states=vpa.accepting_states, - graph=vpa.graph.copy(), - ) - result.validate() - return result - - -def _is_deterministic(vpa: VisiblyPushdownAutomaton) -> bool: - probe = DeterministicVisiblyPushdownAutomaton( - call_alphabet=vpa.call_alphabet, - return_alphabet=vpa.return_alphabet, - internal_alphabet=vpa.internal_alphabet, - stack_alphabet=vpa.stack_alphabet, - bottom_stack_symbol=vpa.bottom_stack_symbol, - initial_state=vpa.initial_state, - accepting_states=vpa.accepting_states, - graph=vpa.graph, - ) - try: - probe._check_determinism() - except NonDeterministicError: - return False - return True - - -class CallDrivenAutomaton(DeterministicVisiblyPushdownAutomaton): - """Deterministic modular VPA whose call target depends only on the call symbol.""" - - modules: dict[Hashable, frozenset[Hashable]] - base_module: Hashable - call_partition: dict[Any, Hashable] - call_entries: dict[Any, Hashable] - - def __init__( - self, - *, - modules: Mapping[Hashable, Iterable[Hashable]] | None = None, - base_module: Hashable = 0, - call_partition: Mapping[Any, Hashable] | None = None, - call_entries: Mapping[Any, Hashable] | None = None, - **kwargs: Any, - ) -> None: - super().__init__(**kwargs) - self.modules = _normalize_modules(modules) - self.base_module = base_module - self.call_partition = dict(call_partition or {}) - self.call_entries = dict(call_entries or {}) - - def validate(self) -> None: - super().validate() - state_modules = self._validate_modules() - self._validate_call_partition() - self._validate_internal_transitions_stay_in_module(state_modules) - self._validate_call_driven_transitions() - - def call_entry_map(self) -> dict[Any, Hashable]: - """Return configured or inferred entries for each call symbol with transitions.""" - entries = dict(self.call_entries) - for transition in self.transitions(): - if transition.data.get(ATTR_KIND) != KIND_CALL: - continue - symbol = transition.data.get(ATTR_SYMBOL) - if symbol is None: - continue - existing = entries.get(symbol, _MISSING) - if existing is _MISSING: - entries[symbol] = transition.target - elif existing != transition.target: - raise NonDeterministicError(f"call target for {symbol!r} depends on source state") - return entries - - @classmethod - def minimize( - cls, - vpa: VisiblyPushdownAutomaton, - *, - modules: Mapping[Hashable, Iterable[Hashable]] | None = None, - call_partition: Mapping[Any, Hashable] | None = None, - base_module: Hashable | None = None, - call_entries: Mapping[Any, Hashable] | None = None, - ) -> CallDrivenAutomaton: - """Return the module-aware deterministic quotient as a CDA.""" - return _minimize_modular_vpa( - cls, - vpa, - modules=modules, - call_partition=call_partition, - base_module=base_module, - call_entries=call_entries, - entry_states=None, - form="cda", - ) - - def _validate_modules(self) -> dict[Hashable, Hashable]: - self._require(bool(self.modules), "modular VPA requires modules") - self._require(self.base_module in self.modules, "base_module must be present in modules") - all_states = set(self.states()) - seen: dict[Hashable, Hashable] = {} - for module, states in self.modules.items(): - self._require(bool(states), f"module {module!r} must contain at least one state") - for state in states: - self._require(state in all_states, f"module {module!r} contains unknown state {state!r}") - self._require(state not in seen, f"state {state!r} appears in multiple modules") - seen[state] = module - self._require(set(seen) == all_states, "modules must cover exactly the VPA states") - if self.initial_state is not None: - self._require( - self.initial_state in self.modules[self.base_module], - "initial_state must lie in the base module", - ) - return seen - - def _validate_call_partition(self) -> None: - missing = self.call_alphabet - set(self.call_partition) - extra = set(self.call_partition) - self.call_alphabet - self._require(not missing, f"call_partition missing calls {sorted(missing, key=repr)!r}") - self._require(not extra, f"call_partition contains non-call symbols {sorted(extra, key=repr)!r}") - for symbol, module in self.call_partition.items(): - self._require(module in self.modules, f"call {symbol!r} targets unknown module {module!r}") - - def _validate_internal_transitions_stay_in_module(self, state_modules: Mapping[Hashable, Hashable]) -> None: - for transition in self.transitions(): - if transition.data.get(ATTR_KIND) == KIND_INTERNAL: - self._require( - state_modules[transition.source] == state_modules[transition.target], - "internal transitions must stay inside one module", - ) - - def _validate_call_driven_transitions(self) -> None: - entries = self.call_entry_map() - for symbol, target in entries.items(): - module = self.call_partition[symbol] - self._require(target in self.modules[module], f"call {symbol!r} entry is not in module {module!r}") - for transition in self.transitions(): - if transition.data.get(ATTR_KIND) != KIND_CALL: - continue - symbol = transition.data.get(ATTR_SYMBOL) - if symbol is None: - continue - self._require( - transition.target == entries[symbol], - f"call target for {symbol!r} must be independent of source state", - ) - - -class MultipleEntryVisiblyPushdownAutomaton(CallDrivenAutomaton): - """Modular VPA with multiple module entries and source-determined call pushes.""" - - entry_states: dict[Hashable, frozenset[Hashable]] - - def __init__( - self, - *, - entry_states: Mapping[Hashable, Iterable[Hashable] | Hashable] | None = None, - **kwargs: Any, - ) -> None: - super().__init__(**kwargs) - self.entry_states = _normalize_multi_entries(entry_states) - - def validate(self) -> None: - super().validate() - self._validate_entry_states() - self._validate_call_targets_are_entries() - self._validate_source_determined_pushes() - - @classmethod - def minimize( - cls, - vpa: VisiblyPushdownAutomaton, - *, - modules: Mapping[Hashable, Iterable[Hashable]] | None = None, - call_partition: Mapping[Any, Hashable] | None = None, - base_module: Hashable | None = None, - entry_states: Mapping[Hashable, Iterable[Hashable] | Hashable] | None = None, - call_entries: Mapping[Any, Hashable] | None = None, - ) -> MultipleEntryVisiblyPushdownAutomaton: - """Return the module-aware deterministic quotient as an MEVPA.""" - return _minimize_modular_vpa( - cls, - vpa, - modules=modules, - call_partition=call_partition, - base_module=base_module, - call_entries=call_entries, - entry_states=entry_states, - form="mevpa", - ) - - def _validate_entry_states(self) -> None: - self._require(bool(self.entry_states), "MEVPA requires entry_states") - for module, states in self.entry_states.items(): - self._require(module in self.modules, f"entry_states has unknown module {module!r}") - self._require(bool(states), f"module {module!r} must have at least one entry") - for state in states: - self._require(state in self.modules[module], f"entry {state!r} is not in module {module!r}") - - def _validate_call_targets_are_entries(self) -> None: - entries = self.call_entry_map() - for symbol, target in entries.items(): - module = self.call_partition[symbol] - self._require( - target in self.entry_states.get(module, frozenset()), - f"call {symbol!r} must enter one of module {module!r}'s entries", - ) - - def _validate_source_determined_pushes(self) -> None: - pushed_by_source: dict[Hashable, Any] = {} - for transition in self.transitions(): - if transition.data.get(ATTR_KIND) != KIND_CALL: - continue - pushed = transition.data.get(ATTR_STACK_SYMBOL) - existing = pushed_by_source.get(transition.source, _MISSING) - if existing is _MISSING: - pushed_by_source[transition.source] = pushed - else: - self._require(existing == pushed, "MEVPA call push must depend only on the source state") - - -class SingleEntryVisiblyPushdownAutomaton(CallDrivenAutomaton): - """Modular VPA with one distinguished entry per non-base module.""" - - entry_states: dict[Hashable, Hashable] - - def __init__( - self, - *, - entry_states: Mapping[Hashable, Hashable] | None = None, - **kwargs: Any, - ) -> None: - super().__init__(**kwargs) - self.entry_states = dict(entry_states or {}) - - def validate(self) -> None: - super().validate() - self._validate_single_entries() - self._validate_single_entry_calls() - - @classmethod - def minimize( - cls, - vpa: VisiblyPushdownAutomaton, - *, - call_partition: Mapping[Any, Hashable] | None = None, - modules: Mapping[Hashable, Iterable[Hashable]] | None = None, - base_module: Hashable | None = None, - entry_states: Mapping[Hashable, Hashable] | None = None, - call_entries: Mapping[Any, Hashable] | None = None, - ) -> SingleEntryVisiblyPushdownAutomaton: - """Return the module-aware deterministic quotient as an SEVPA. - - A fixed call partition and module structure are required. General VPA - minimization is intentionally not attempted here. - """ - return _minimize_modular_vpa( - cls, - vpa, - modules=modules, - call_partition=call_partition, - base_module=base_module, - call_entries=call_entries, - entry_states=entry_states, - form="sevpa", - ) - - def _validate_single_entries(self) -> None: - missing = set(self.modules) - {self.base_module} - set(self.entry_states) - self._require(not missing, f"SEVPA missing entries for modules {sorted(missing, key=repr)!r}") - for module, state in self.entry_states.items(): - self._require(module in self.modules, f"entry_states has unknown module {module!r}") - self._require(state in self.modules[module], f"entry {state!r} is not in module {module!r}") - - def _validate_single_entry_calls(self) -> None: - for transition in self.transitions(): - if transition.data.get(ATTR_KIND) != KIND_CALL: - continue - symbol = transition.data.get(ATTR_SYMBOL) - if symbol is None: - continue - module = self.call_partition[symbol] - expected_entry = self.entry_states.get(module) - self._require( - transition.target == expected_entry, - f"SEVPA call {symbol!r} must enter module {module!r}'s single entry", - ) - self._require( - transition.data.get(ATTR_STACK_SYMBOL) == (transition.source, symbol), - "SEVPA call stack symbols must be (caller_state, call_symbol)", - ) - - -class CanonicalVisiblyPushdownAutomaton(DeterministicVisiblyPushdownAutomaton): - """Canonical VPA built from the finite Myhill-Nerode summary algebra.""" - - summary_representatives: dict[Hashable, tuple[int | None, ...]] - - def __init__( - self, - *, - summary_representatives: Mapping[Hashable, tuple[int | None, ...]] | None = None, - **kwargs: Any, - ) -> None: - super().__init__(**kwargs) - self.summary_representatives = dict(summary_representatives or {}) - - @classmethod - def from_vpa(cls, vpa: VisiblyPushdownAutomaton) -> CanonicalVisiblyPushdownAutomaton: - """Build the Myhill-Nerode canonical deterministic VPA of a well-matched language. - - States are classes of the summary algebra of well-matched factors, so - the form is canonical for well-matched languages. Raises - :class:`~sofic.exceptions.NonWellMatchedLanguageError` when ``vpa`` - accepts a word with a pending call or return; general VPLs have no - unique minimal deterministic VPA :cite:`AlurKumarMadhusudanViswanathan2005`. - With an empty call alphabet this is the minimal DFA. - """ - if vpa.has_unmatched_word(): - raise NonWellMatchedLanguageError( - "the canonical VPA is defined for well-matched languages; this one accepts a pending call or return" - ) - det = DeterministicVisiblyPushdownAutomaton.from_vpa(vpa) - return _SummaryAlgebra(det).to_canonical_vpa(cls) - - @classmethod - def minimize(cls, vpa: VisiblyPushdownAutomaton) -> CanonicalVisiblyPushdownAutomaton: - """Alias for :meth:`from_vpa`.""" - return cls.from_vpa(vpa) - - -def _normalize_modules(modules: Mapping[Hashable, Iterable[Hashable]] | None) -> dict[Hashable, frozenset[Hashable]]: - if modules is None: - return {} - return {module: frozenset(states) for module, states in modules.items()} - - -def _normalize_multi_entries( - entry_states: Mapping[Hashable, Iterable[Hashable] | Hashable] | None, -) -> dict[Hashable, frozenset[Hashable]]: - if entry_states is None: - return {} - result: dict[Hashable, frozenset[Hashable]] = {} - for module, states in entry_states.items(): - if isinstance(states, frozenset | set | list): - result[module] = frozenset(states) - else: - result[module] = frozenset({states}) - return result - - -def _metadata_or_argument(vpa: VisiblyPushdownAutomaton, name: str, value: Any, default: Any) -> Any: - if value is not None: - return value - return getattr(vpa, name, default) - - -def _deterministic_view(vpa: VisiblyPushdownAutomaton) -> DeterministicVisiblyPushdownAutomaton: - return DeterministicVisiblyPushdownAutomaton.from_vpa(vpa) - - -def _minimize_modular_vpa( - target_cls: type[CallDrivenAutomaton], - vpa: VisiblyPushdownAutomaton, - *, - modules: Mapping[Hashable, Iterable[Hashable]] | None, - call_partition: Mapping[Any, Hashable] | None, - base_module: Hashable | None, - call_entries: Mapping[Any, Hashable] | None, - entry_states: Mapping[Hashable, Any] | None, - form: str, -) -> Any: - modules = _metadata_or_argument(vpa, "modules", modules, None) - call_partition = _metadata_or_argument(vpa, "call_partition", call_partition, None) - base_module = _metadata_or_argument(vpa, "base_module", base_module, 0) - call_entries = _metadata_or_argument(vpa, "call_entries", call_entries, None) - entry_states = _metadata_or_argument(vpa, "entry_states", entry_states, None) - - if modules is None: - convert = vc.to_multiple_entry if form == "mevpa" else vc.to_single_entry - converted = convert(vpa, call_partition) - return _minimize_modular_vpa( - target_cls, - converted, - modules=converted.modules, - call_partition=converted.call_partition, - base_module=converted.base_module, - call_entries=converted.call_entries, - entry_states=getattr(converted, "entry_states", None) if form != "cda" else None, - form=form, - ) - if call_partition is None: - raise ValueError("a call_partition is required when modules are given") - - det = _deterministic_view(vpa) - modules = _normalize_modules(modules) - call_partition = dict(call_partition) - call_entries = dict(call_entries or _infer_call_entries(det, call_partition)) - entry_states = _infer_entry_states(form, modules, base_module, call_partition, call_entries, entry_states, det) - - source = _construct_modular_view( - target_cls, - det, - modules=modules, - base_module=base_module, - call_partition=call_partition, - call_entries=call_entries, - entry_states=entry_states, - ) - source.validate() - - partition = _refine_modular_partition(source, modules) - return _quotient_modular_vpa( - target_cls, - source, - partition, - modules=modules, - base_module=base_module, - call_partition=call_partition, - call_entries=call_entries, - entry_states=entry_states, - form=form, - ) - - -def _construct_modular_view( - target_cls: type[CallDrivenAutomaton], - det: DeterministicVisiblyPushdownAutomaton, - *, - modules: Mapping[Hashable, frozenset[Hashable]], - base_module: Hashable, - call_partition: Mapping[Any, Hashable], - call_entries: Mapping[Any, Hashable], - entry_states: Mapping[Hashable, Any], -) -> CallDrivenAutomaton: - kwargs = { - "input_alphabet": det.input_alphabet, - "call_alphabet": det.call_alphabet, - "return_alphabet": det.return_alphabet, - "internal_alphabet": det.internal_alphabet, - "stack_alphabet": det.stack_alphabet, - "bottom_stack_symbol": det.bottom_stack_symbol, - "initial_state": det.initial_state, - "accepting_states": det.accepting_states, - "graph": det.graph.copy(), - "modules": modules, - "base_module": base_module, - "call_partition": call_partition, - "call_entries": call_entries, - } - if issubclass(target_cls, (SingleEntryVisiblyPushdownAutomaton, MultipleEntryVisiblyPushdownAutomaton)): - kwargs["entry_states"] = entry_states - return target_cls(**kwargs) - - -def _infer_call_entries( - det: DeterministicVisiblyPushdownAutomaton, - call_partition: Mapping[Any, Hashable], -) -> dict[Any, Hashable]: - entries: dict[Any, Hashable] = {} - for (_source, symbol), (target, _stack) in det.call_transition_map().items(): - if symbol not in call_partition: - raise ValueError(f"call_partition missing call symbol {symbol!r}") - existing = entries.get(symbol, _MISSING) - if existing is _MISSING: - entries[symbol] = target - elif existing != target: - raise ValueError( - f"call target for {symbol!r} depends on the source state; omit modules to convert the VPA first" - ) - return entries - - -def _infer_entry_states( - form: str, - modules: Mapping[Hashable, frozenset[Hashable]], - base_module: Hashable, - call_partition: Mapping[Any, Hashable], - call_entries: Mapping[Any, Hashable], - entry_states: Mapping[Hashable, Any] | None, - det: DeterministicVisiblyPushdownAutomaton, -) -> dict[Hashable, Any]: - if entry_states is not None: - if form == "mevpa": - return _normalize_multi_entries(entry_states) - return dict(entry_states) - - by_module: dict[Hashable, set[Hashable]] = {module: set() for module in modules} - if det.initial_state is not None and base_module in by_module: - by_module[base_module].add(det.initial_state) - for symbol, entry in call_entries.items(): - by_module[call_partition[symbol]].add(entry) - - if form == "mevpa": - return {module: frozenset(states) for module, states in by_module.items() if states} - if form == "sevpa": - result: dict[Hashable, Hashable] = {} - for module, states in by_module.items(): - if module == base_module: - continue - if len(states) != 1: - raise ValueError( - "SEVPA minimization requires one entry per non-base module; omit modules to convert the VPA first" - ) - result[module] = next(iter(states)) - return result - return {} - - -def _refine_modular_partition( - vpa: CallDrivenAutomaton, - modules: Mapping[Hashable, frozenset[Hashable]], -) -> list[frozenset[Hashable]]: - partition: list[frozenset[Hashable]] = [] - for _module, states in sorted(modules.items(), key=lambda item: repr(item[0])): - accepting = frozenset(states & vpa.accepting_states) - rejecting = frozenset(states - vpa.accepting_states) - if accepting: - partition.append(accepting) - if rejecting: - partition.append(rejecting) - - changed = True - while changed: - changed = False - block_of = _block_map(partition) - context_groups = _stack_context_groups(vpa, block_of) - new_partition: list[frozenset[Hashable]] = [] - for block in partition: - pieces: dict[tuple[Any, ...], set[Hashable]] = {} - for state in block: - signature = _modular_state_signature(vpa, state, block_of, context_groups) - pieces.setdefault(signature, set()).add(state) - if len(pieces) > 1: - changed = True - new_partition.extend(frozenset(piece) for piece in pieces.values()) - partition = new_partition - return partition - - -def _block_map(partition: Sequence[frozenset[Hashable]]) -> dict[Hashable, frozenset[Hashable]]: - return {state: block for block in partition for state in block} - - -def _stack_context_groups( - vpa: DeterministicVisiblyPushdownAutomaton, - block_of: Mapping[Hashable, Hashable], -) -> list[tuple[Any, tuple[Any, ...]]]: - stack_symbols = {stack for _key, (_target, stack) in vpa.call_transition_map().items()} - if vpa.bottom_stack_symbol is not None: - stack_symbols.add(vpa.bottom_stack_symbol) - grouped: dict[Any, set[Any]] = {} - for stack_symbol in stack_symbols: - canonical = _canonical_stack_symbol(stack_symbol, block_of) - grouped.setdefault(canonical, set()).add(stack_symbol) - return [(canonical, tuple(sorted(actuals, key=repr))) for canonical, actuals in sorted(grouped.items(), key=repr)] - - -def _modular_state_signature( - vpa: CallDrivenAutomaton, - state: Hashable, - block_of: Mapping[Hashable, Hashable], - context_groups: Sequence[tuple[Any, tuple[Any, ...]]], -) -> tuple[Any, ...]: - state_modules = {state: module for module, states in vpa.modules.items() for state in states} - call_map = vpa.call_transition_map() - internal_map = vpa.internal_transition_map() - return_map = vpa.return_transition_map() - - internal = tuple( - ( - symbol, - None if (target := internal_map.get((state, symbol))) is None else block_of[target], - ) - for symbol in sorted(vpa.internal_alphabet, key=repr) - ) - calls = tuple( - ( - symbol, - None - if (call := call_map.get((state, symbol))) is None - else (block_of[call[0]], _canonical_stack_symbol(call[1], block_of)), - ) - for symbol in sorted(vpa.call_alphabet, key=repr) - ) - returns = [] - for symbol in sorted(vpa.return_alphabet, key=repr): - for canonical_stack, actual_stacks in context_groups: - targets = [] - for stack_symbol in actual_stacks: - target = return_map.get((state, symbol, stack_symbol), return_map.get((state, symbol, None))) - targets.append(None if target is None else block_of[target]) - returns.append((symbol, canonical_stack, tuple(sorted(set(targets), key=repr)))) - - return ( - state_modules[state], - state in vpa.accepting_states, - internal, - calls, - tuple(returns), - ) - - -def _quotient_modular_vpa( - target_cls: type[CallDrivenAutomaton], - source: CallDrivenAutomaton, - partition: Sequence[frozenset[Hashable]], - *, - modules: Mapping[Hashable, frozenset[Hashable]], - base_module: Hashable, - call_partition: Mapping[Any, Hashable], - call_entries: Mapping[Any, Hashable], - entry_states: Mapping[Hashable, Any], - form: str, -) -> Any: - block_of = _block_map(partition) - blocks = list(partition) - stack_alphabet: set[Any] = set() - if source.bottom_stack_symbol is not None: - stack_alphabet.add(source.bottom_stack_symbol) - transitions: set[tuple[Hashable, Hashable, Any, Any, Any | None]] = set() - stack_symbol_rewrite: dict[Any, set[Any]] = {} - - for transition in source.transitions(): - if transition.data.get(ATTR_KIND) != KIND_CALL: - continue - symbol = transition.data.get(ATTR_SYMBOL) - stack_symbol = transition.data.get(ATTR_STACK_SYMBOL) - quotient_stack_symbol = _quotient_call_stack_symbol( - form, - block_of[transition.source], - symbol, - stack_symbol, - block_of, - ) - stack_symbol_rewrite.setdefault(stack_symbol, set()).add(quotient_stack_symbol) - - for block in blocks: - representative = min(block, key=repr) - source_block = block_of[representative] - for transition in source.graph.out_transitions(representative): - kind = transition.data.get(ATTR_KIND) - symbol = transition.data.get(ATTR_SYMBOL) - target_block = block_of[transition.target] - stack_symbol = transition.data.get(ATTR_STACK_SYMBOL) - if kind == KIND_CALL: - stack_symbol = _quotient_call_stack_symbol(form, source_block, symbol, stack_symbol, block_of) - stack_alphabet.add(stack_symbol) - elif kind == KIND_RETURN and stack_symbol is not None: - rewritten = stack_symbol_rewrite.get(stack_symbol, {_canonical_stack_symbol(stack_symbol, block_of)}) - for rewritten_stack_symbol in rewritten: - stack_alphabet.add(rewritten_stack_symbol) - transitions.add((source_block, target_block, kind, symbol, rewritten_stack_symbol)) - continue - transitions.add((source_block, target_block, kind, symbol, stack_symbol)) - - quotient_modules = { - module: frozenset(block for block in blocks if block & states) - for module, states in modules.items() - if any(block & states for block in blocks) - } - quotient_call_entries = {symbol: block_of[state] for symbol, state in call_entries.items() if state in block_of} - quotient_entry_states = _quotient_entry_states(form, entry_states, block_of) - - kwargs = { - "input_alphabet": source.input_alphabet, - "call_alphabet": source.call_alphabet, - "return_alphabet": source.return_alphabet, - "internal_alphabet": source.internal_alphabet, - "stack_alphabet": frozenset(stack_alphabet), - "bottom_stack_symbol": source.bottom_stack_symbol, - "initial_state": block_of[source.initial_state], - "accepting_states": frozenset(block for block in blocks if block & source.accepting_states), - "modules": quotient_modules, - "base_module": base_module, - "call_partition": dict(call_partition), - "call_entries": quotient_call_entries, - } - if issubclass(target_cls, (SingleEntryVisiblyPushdownAutomaton, MultipleEntryVisiblyPushdownAutomaton)): - kwargs["entry_states"] = quotient_entry_states - - result = target_cls(**kwargs) - for block in blocks: - result.graph.add_state(block) - for source_block, target_block, kind, symbol, stack_symbol in sorted(transitions, key=repr): - if kind == KIND_CALL: - result.add_call_transition(source_block, target_block, symbol, stack_symbol) - elif kind == KIND_RETURN: - result.add_return_transition(source_block, target_block, symbol, stack_symbol) - elif kind == KIND_INTERNAL: - result.add_internal_transition(source_block, target_block, symbol) - result.validate() - return result - - -def _quotient_call_stack_symbol( - form: str, - source_block: Hashable, - symbol: Any, - stack_symbol: Any, - block_of: Mapping[Hashable, Hashable], -) -> Any: - if form == "sevpa": - return (source_block, symbol) - if form == "mevpa": - return source_block - return _canonical_stack_symbol(stack_symbol, block_of) - - -def _quotient_entry_states( - form: str, - entry_states: Mapping[Hashable, Any], - block_of: Mapping[Hashable, Hashable], -) -> dict[Hashable, Any]: - if form == "mevpa": - normalized = _normalize_multi_entries(entry_states) - return { - module: frozenset(block_of[state] for state in states if state in block_of) - for module, states in normalized.items() - } - if form == "sevpa": - return {module: block_of[state] for module, state in entry_states.items() if state in block_of} - return {} - - -def _canonical_stack_symbol(stack_symbol: Any, block_of: Mapping[Hashable, Hashable]) -> Any: - if stack_symbol in block_of: - return block_of[stack_symbol] - if isinstance(stack_symbol, tuple): - return tuple(_canonical_stack_symbol(part, block_of) for part in stack_symbol) - return stack_symbol - - -class _SummaryAlgebra: - """Finite algebra of well-matched summaries with top-level and nested congruences. - - A summary maps each source state to the state reached along a well-matched - word (``None`` when the run dies). Top-level summaries are identified when - every well-matched continuation accepts both or neither; nested summaries - (inside a pending call) are identified when, for every enclosing context, - returning from them lands in the same class. Both partitions are refined - together until internal steps, calls, and returns are well defined on - classes, which makes the quotient canonical. - """ - - def __init__(self, vpa: DeterministicVisiblyPushdownAutomaton) -> None: - self.vpa = vpa - self.state_order = tuple(sorted(vpa.states(), key=repr)) - state_index = {state: index for index, state in enumerate(self.state_order)} - self.internal_summaries = { - symbol: _internal_summary(vpa, self.state_order, state_index, symbol) - for symbol in sorted(vpa.internal_alphabet, key=repr) - } - self.identity = tuple(range(len(self.state_order))) - self.summaries = sorted( - _close_summary_algebra(vpa, self.state_order, state_index, self.identity, self.internal_summaries), - key=repr, - ) - self.calls = sorted(vpa.call_alphabet, key=repr) - self.returns = sorted(vpa.return_alphabet, key=repr) - self._wrap_cache: dict[tuple, tuple] = {} - self.top, self.nested = self._refine() - - def wrap(self, inner: tuple, call: Any, ret: Any) -> tuple: - key = (inner, call, ret) - if key not in self._wrap_cache: - self._wrap_cache[key] = _wrap_summary(self.vpa, self.state_order, inner, call, ret) - return self._wrap_cache[key] - - def accepts(self, summary: tuple) -> bool: - return _summary_accepts(self.vpa, self.state_order, summary) - - def _refine(self) -> tuple[dict[tuple, int], dict[tuple, int]]: - top = {s: int(self.accepts(s)) for s in self.summaries} - nested = dict.fromkeys(self.summaries, 0) - internals = list(self.internal_summaries.values()) - pairs = [(c, r) for c in self.calls for r in self.returns] - while True: - top_signature = { - s: ( - top[s], - tuple(top[_compose_summary(s, a)] for a in internals), - tuple(top[_compose_summary(s, self.wrap(x, c, r))] for x in self.summaries for c, r in pairs), - ) - for s in self.summaries - } - nested_signature = { - s: ( - nested[s], - tuple(nested[_compose_summary(s, a)] for a in internals), - tuple(nested[_compose_summary(s, self.wrap(x, c, r))] for x in self.summaries for c, r in pairs), - tuple( - (top[_compose_summary(o, self.wrap(s, c, r))], nested[_compose_summary(o, self.wrap(s, c, r))]) - for o in self.summaries - for c, r in pairs - ), - ) - for s in self.summaries - } - new_top = _number_blocks(top_signature, first=top_signature[self.identity]) - new_nested = _number_blocks(nested_signature, first=nested_signature[self.identity]) - stable = len(set(new_top.values())) == len(set(top.values())) and len(set(new_nested.values())) == len( - set(nested.values()) - ) - top, nested = new_top, new_nested - if stable: - return top, nested - - def to_canonical_vpa(self, cls: type[CanonicalVisiblyPushdownAutomaton]) -> CanonicalVisiblyPushdownAutomaton: - representative: dict[Hashable, tuple] = {} - for s in self.summaries: - representative.setdefault(self.top[s], s) - representative.setdefault(("nested", self.nested[s]), s) - - def label(summary: tuple, nested: bool) -> Hashable: - return ("nested", self.nested[summary]) if nested else self.top[summary] - - initial = label(self.identity, False) - entry = label(self.identity, True) - reachable: set[Hashable] = set() - internals, calls, returns = set(), set(), set() - pushed: set[tuple[Hashable, Any]] = set() - frontier = [initial] - while frontier: - while frontier: - state = frontier.pop() - if state in reachable: - continue - reachable.add(state) - summary = representative[state] - for symbol, step in self.internal_summaries.items(): - target = label(_compose_summary(summary, step), isinstance(state, tuple)) - internals.add((state, symbol, target)) - frontier.append(target) - for call in self.calls: - calls.add((state, call, entry, (state, call))) - pushed.add((state, call)) - frontier.append(entry) - # Returns pop a pushed (outer, call) from a nested state; both sets - # grow together, so repeat until no new state appears. - for state in [s for s in reachable if isinstance(s, tuple)]: - for outer, call in list(pushed): - for ret in self.returns: - combined = _compose_summary(representative[outer], self.wrap(representative[state], call, ret)) - target = label(combined, isinstance(outer, tuple)) - returns.add((state, ret, (outer, call), target)) - if target not in reachable: - frontier.append(target) - - result = cls( - input_alphabet=self.vpa.input_alphabet, - call_alphabet=self.vpa.call_alphabet, - return_alphabet=self.vpa.return_alphabet, - internal_alphabet=self.vpa.internal_alphabet, - stack_alphabet=frozenset(pushed), - bottom_stack_symbol=None, - initial_state=initial, - accepting_states=frozenset( - s for s in reachable if not isinstance(s, tuple) and self.accepts(representative[s]) - ), - summary_representatives={s: representative[s] for s in reachable}, - ) - for state in sorted(reachable, key=repr): - result.graph.add_state(state) - for source, symbol, target in sorted(internals, key=repr): - result.add_internal_transition(source, target, symbol) - for source, symbol, target, push in sorted(calls, key=repr): - result.add_call_transition(source, target, symbol, push) - for source, symbol, guard, target in sorted(returns, key=repr): - result.add_return_transition(source, target, symbol, guard) - result.validate() - return result - - -def _number_blocks(signatures: Mapping[tuple, Any], *, first: Any) -> dict[tuple, int]: - ordered = sorted(set(signatures.values()), key=lambda sig: (sig != first, repr(sig))) - index = {signature: position for position, signature in enumerate(ordered)} - return {summary: index[signature] for summary, signature in signatures.items()} - - -def _internal_summary( - vpa: DeterministicVisiblyPushdownAutomaton, - state_order: Sequence[Hashable], - state_index: Mapping[Hashable, int], - symbol: Any, -) -> tuple[int | None, ...]: - transitions = vpa.internal_transition_map() - summary: list[int | None] = [] - for state in state_order: - target = transitions.get((state, symbol)) - summary.append(None if target is None else state_index[target]) - return tuple(summary) - - -def _close_summary_algebra( - vpa: DeterministicVisiblyPushdownAutomaton, - state_order: Sequence[Hashable], - state_index: Mapping[Hashable, int], - identity: tuple[int | None, ...], - internal_summaries: Mapping[Any, tuple[int | None, ...]], -) -> set[tuple[int | None, ...]]: - summaries = {identity, *internal_summaries.values()} - queue: deque[tuple[int | None, ...]] = deque(sorted(summaries, key=repr)) - while queue: - summary = queue.popleft() - current = list(summaries) - candidates: list[tuple[int | None, ...]] = [] - for other in current: - candidates.append(_compose_summary(summary, other)) - candidates.append(_compose_summary(other, summary)) - for call_symbol in sorted(vpa.call_alphabet, key=repr): - for return_symbol in sorted(vpa.return_alphabet, key=repr): - candidates.append(_wrap_summary(vpa, state_order, summary, call_symbol, return_symbol)) - for candidate in candidates: - if candidate not in summaries: - summaries.add(candidate) - queue.append(candidate) - return summaries - - -def _compose_summary( - first: tuple[int | None, ...], - second: tuple[int | None, ...], -) -> tuple[int | None, ...]: - return tuple(None if state is None else second[state] for state in first) - - -def _wrap_summary( - vpa: DeterministicVisiblyPushdownAutomaton, - state_order: Sequence[Hashable], - inner: tuple[int | None, ...], - call_symbol: Any, - return_symbol: Any, -) -> tuple[int | None, ...]: - state_index = {state: index for index, state in enumerate(state_order)} - call_map = vpa.call_transition_map() - result: list[int | None] = [] - for state in state_order: - call = call_map.get((state, call_symbol)) - if call is None: - result.append(None) - continue - call_target, stack_symbol = call - inner_target_index = inner[state_index[call_target]] - if inner_target_index is None: - result.append(None) - continue - return_target = vpa.return_successor(state_order[inner_target_index], return_symbol, stack_symbol) - result.append(None if return_target is None else state_index[return_target]) - return tuple(result) - - -def _summary_accepts( - vpa: DeterministicVisiblyPushdownAutomaton, - state_order: Sequence[Hashable], - summary: tuple[int | None, ...], -) -> bool: - if vpa.initial_state is None: - return False - initial_index = {state: index for index, state in enumerate(state_order)}[vpa.initial_state] - target = summary[initial_index] - return target is not None and state_order[target] in vpa.accepting_states diff --git a/sofic/automata/vpa/__init__.py b/sofic/automata/vpa/__init__.py new file mode 100644 index 0000000..3e7112b --- /dev/null +++ b/sofic/automata/vpa/__init__.py @@ -0,0 +1,26 @@ +"""Visibly pushdown automata, their constructions, and simulation.""" + +from sofic.automata.vpa.base import VisiblyPushdownAutomaton +from sofic.automata.vpa.canonical import CanonicalVisiblyPushdownAutomaton +from sofic.automata.vpa.deterministic import DeterministicVisiblyPushdownAutomaton +from sofic.automata.vpa.modular import ( + CallDrivenAutomaton, + MultipleEntryVisiblyPushdownAutomaton, + SingleEntryVisiblyPushdownAutomaton, +) +from sofic.automata.vpa.operations import BOTTOM, NormalVPA, to_multiple_entry, to_single_entry +from sofic.automata.vpa.simulation import recognizes_vpa + +__all__ = [ + "BOTTOM", + "CallDrivenAutomaton", + "CanonicalVisiblyPushdownAutomaton", + "DeterministicVisiblyPushdownAutomaton", + "MultipleEntryVisiblyPushdownAutomaton", + "NormalVPA", + "SingleEntryVisiblyPushdownAutomaton", + "VisiblyPushdownAutomaton", + "recognizes_vpa", + "to_multiple_entry", + "to_single_entry", +] diff --git a/sofic/automata/vpa/base.py b/sofic/automata/vpa/base.py new file mode 100644 index 0000000..ca20f65 --- /dev/null +++ b/sofic/automata/vpa/base.py @@ -0,0 +1,251 @@ +"""Visibly pushdown automata.""" + +from __future__ import annotations + +from collections.abc import Hashable, Sequence +from typing import TYPE_CHECKING, Any + +from sofic.automata.vpa import operations as vc +from sofic.base import StateMachine +from sofic.exceptions import NonDeterministicError +from sofic.graph import ( + ATTR_KIND, + ATTR_STACK_SYMBOL, + ATTR_SYMBOL, + KIND_CALL, + KIND_INTERNAL, + KIND_RETURN, +) + +if TYPE_CHECKING: + from sofic.automata.vpa.deterministic import DeterministicVisiblyPushdownAutomaton + +_MISSING = object() + + +class VisiblyPushdownAutomaton(StateMachine): + """Standard 1-stack VPA with call / return / internal input partition.""" + + input_alphabet: frozenset[Any] + call_alphabet: frozenset[Any] + return_alphabet: frozenset[Any] + internal_alphabet: frozenset[Any] + stack_alphabet: frozenset[Any] + bottom_stack_symbol: Any | None + initial_state: Hashable | None + accepting_states: frozenset[Hashable] + + def __init__( + self, + input_alphabet: frozenset[Any] | None = None, + call_alphabet: frozenset[Any] | None = None, + return_alphabet: frozenset[Any] | None = None, + internal_alphabet: frozenset[Any] | None = None, + stack_alphabet: frozenset[Any] | None = None, + bottom_stack_symbol: Any | None = None, + initial_state: Hashable | None = None, + accepting_states: frozenset[Hashable] | None = None, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + self.call_alphabet = call_alphabet if call_alphabet is not None else frozenset() + self.return_alphabet = return_alphabet if return_alphabet is not None else frozenset() + self.internal_alphabet = internal_alphabet if internal_alphabet is not None else frozenset() + self.stack_alphabet = stack_alphabet if stack_alphabet is not None else frozenset() + self.bottom_stack_symbol = bottom_stack_symbol + self.input_alphabet = ( + input_alphabet + if input_alphabet is not None + else self.call_alphabet | self.return_alphabet | self.internal_alphabet + ) + self.initial_state = initial_state + self.accepting_states = accepting_states if accepting_states is not None else frozenset() + + def validate(self) -> None: + partition = self.call_alphabet | self.return_alphabet | self.internal_alphabet + self._require( + len(self.call_alphabet) + len(self.return_alphabet) + len(self.internal_alphabet) == len(partition), + "call, return, and internal alphabets must be disjoint", + ) + self._require(partition == self.input_alphabet, "input alphabet must equal partition of call/return/internal") + if self.bottom_stack_symbol is not None: + self._require( + self.bottom_stack_symbol in self.stack_alphabet, "bottom_stack_symbol must be in stack alphabet" + ) + if self.initial_state is not None: + self._require(self.graph.has_state(self.initial_state), "missing initial state") + for state in self.accepting_states: + self._require(self.graph.has_state(state), f"missing accepting state {state!r}") + for transition in self.transitions(): + kind = transition.data.get(ATTR_KIND) + symbol = transition.data.get(ATTR_SYMBOL) + self._require(kind in {KIND_CALL, KIND_RETURN, KIND_INTERNAL}, f"invalid VPA kind {kind!r}") + if symbol is not None: + if kind == KIND_CALL: + self._require(symbol in self.call_alphabet, f"{symbol!r} not in call alphabet") + stack_sym = transition.data.get(ATTR_STACK_SYMBOL) + self._require(stack_sym in self.stack_alphabet, "call edge requires stack_symbol in stack alphabet") + self._require( + stack_sym != self.bottom_stack_symbol, + "call edge cannot push the bottom_stack_symbol", + ) + elif kind == KIND_RETURN: + self._require(symbol in self.return_alphabet, f"{symbol!r} not in return alphabet") + stack_sym = transition.data.get(ATTR_STACK_SYMBOL) + if stack_sym is not None: + self._require(stack_sym in self.stack_alphabet, "return stack_symbol must be in stack alphabet") + else: + self._require(symbol in self.internal_alphabet, f"{symbol!r} not in internal alphabet") + + def add_call_transition( + self, + source: Hashable, + target: Hashable, + symbol: Any, + stack_symbol: Any, + **attrs: Any, + ) -> int: + """Add a call transition that pushes ``stack_symbol``.""" + data = {**attrs, ATTR_KIND: KIND_CALL, ATTR_SYMBOL: symbol, ATTR_STACK_SYMBOL: stack_symbol} + return self.graph.add_transition(source, target, **data) + + def add_return_transition( + self, + source: Hashable, + target: Hashable, + symbol: Any, + stack_symbol: Any | None = None, + **attrs: Any, + ) -> int: + """Add a return transition. + + If ``stack_symbol`` is omitted, the transition is a wildcard: it fires + on every stack symbol, and also on the empty stack (leaving it empty) + when the VPA has a ``bottom_stack_symbol``. A return guarded by the + bottom symbol fires only on the empty stack. + """ + data = {**attrs, ATTR_KIND: KIND_RETURN, ATTR_SYMBOL: symbol} + if stack_symbol is not None: + data[ATTR_STACK_SYMBOL] = stack_symbol + return self.graph.add_transition(source, target, **data) + + def add_internal_transition(self, source: Hashable, target: Hashable, symbol: Any, **attrs: Any) -> int: + """Add an internal transition.""" + data = {**attrs, ATTR_KIND: KIND_INTERNAL, ATTR_SYMBOL: symbol} + return self.graph.add_transition(source, target, **data) + + def call_transition_map(self) -> dict[tuple[Hashable, Any], tuple[Hashable, Any]]: + """Return deterministic call transitions keyed by ``(state, symbol)``.""" + result: dict[tuple[Hashable, Any], tuple[Hashable, Any]] = {} + for transition in self.transitions(): + if transition.data.get(ATTR_KIND) != KIND_CALL: + continue + symbol = transition.data.get(ATTR_SYMBOL) + if symbol is None: + continue + key = (transition.source, symbol) + value = (transition.target, transition.data.get(ATTR_STACK_SYMBOL)) + if key in result and result[key] != value: + raise NonDeterministicError(f"non-deterministic call transition on {key}") + result[key] = value + return result + + def return_transition_map(self) -> dict[tuple[Hashable, Any, Any | None], Hashable]: + """Return deterministic return transitions keyed by ``(state, symbol, stack_symbol)``.""" + result: dict[tuple[Hashable, Any, Any | None], Hashable] = {} + for transition in self.transitions(): + if transition.data.get(ATTR_KIND) != KIND_RETURN: + continue + symbol = transition.data.get(ATTR_SYMBOL) + if symbol is None: + continue + key = (transition.source, symbol, transition.data.get(ATTR_STACK_SYMBOL)) + value = transition.target + if key in result and result[key] != value: + raise NonDeterministicError(f"non-deterministic return transition on {key}") + result[key] = value + return result + + def internal_transition_map(self) -> dict[tuple[Hashable, Any], Hashable]: + """Return deterministic internal transitions keyed by ``(state, symbol)``.""" + result: dict[tuple[Hashable, Any], Hashable] = {} + for transition in self.transitions(): + if transition.data.get(ATTR_KIND) != KIND_INTERNAL: + continue + symbol = transition.data.get(ATTR_SYMBOL) + if symbol is None: + continue + key = (transition.source, symbol) + value = transition.target + if key in result and result[key] != value: + raise NonDeterministicError(f"non-deterministic internal transition on {key}") + result[key] = value + return result + + def recognizes(self, word: Sequence[Any]) -> bool: + from sofic.automata.vpa.simulation import recognizes_vpa + + return recognizes_vpa(self, word) + + def union(self, other: VisiblyPushdownAutomaton) -> VisiblyPushdownAutomaton: + """Return a VPA recognizing the union with ``other``.""" + return _binary(vc.union, self, other) + + def intersection(self, other: VisiblyPushdownAutomaton) -> VisiblyPushdownAutomaton: + """Return a VPA recognizing the intersection with ``other`` (synchronized product).""" + return _binary(vc.intersection, self, other) + + def complement(self) -> DeterministicVisiblyPushdownAutomaton: + """Return a deterministic VPA recognizing the complement over this visible alphabet.""" + from sofic.automata.vpa.deterministic import DeterministicVisiblyPushdownAutomaton + + return vc.denormalize(vc.complement(vc.normalize(self)), DeterministicVisiblyPushdownAutomaton) + + def difference(self, other: VisiblyPushdownAutomaton) -> VisiblyPushdownAutomaton: + """Return a VPA recognizing this language minus ``other``.""" + return _binary(vc.difference, self, other) + + def concat(self, other: VisiblyPushdownAutomaton) -> VisiblyPushdownAutomaton: + """Return a VPA recognizing concatenation with ``other``. + + Each factor is read from an empty stack of its own: a return in the right + factor that would pop a pending call of the left factor is a pending + return of the right factor. + """ + return _binary(vc.concat, self, other) + + def kleene_star(self) -> VisiblyPushdownAutomaton: + """Return a VPA recognizing the Kleene star; each factor starts from its own empty stack.""" + return vc.denormalize(vc.kleene_star(vc.normalize(self))) + + def determinize(self) -> DeterministicVisiblyPushdownAutomaton: + """Return an equivalent complete deterministic VPA :cite:`AlurMadhusudan2009`.""" + return vc.determinize_vpa(self) + + def is_empty(self) -> bool: + """Return whether the language is empty.""" + return vc.is_empty(vc.normalize(self)) + + def accepted_word(self) -> tuple[Any, ...] | None: + """Return a short accepted word, or ``None`` when the language is empty.""" + return vc.accepted_word(vc.normalize(self)) + + def is_universal(self) -> bool: + """Return whether every word over the visible alphabet is accepted.""" + return vc.is_empty(vc.complement(vc.normalize(self))) + + def includes(self, other: VisiblyPushdownAutomaton) -> bool: + """Return whether ``other``'s language is contained in this one.""" + return vc.is_empty(vc.difference(vc.normalize(other), vc.normalize(self))) + + def equivalent(self, other: VisiblyPushdownAutomaton) -> bool: + """Return whether both VPAs recognize the same language.""" + return self.includes(other) and other.includes(self) + + def has_unmatched_word(self) -> bool: + """Return whether some accepted word has a pending call or a pending return.""" + return vc.has_unmatched_word(vc.normalize(self)) + + +def _binary(operation, left: VisiblyPushdownAutomaton, right: VisiblyPushdownAutomaton) -> VisiblyPushdownAutomaton: + return vc.denormalize(operation(vc.normalize(left), vc.normalize(right))) diff --git a/sofic/automata/vpa/canonical.py b/sofic/automata/vpa/canonical.py new file mode 100644 index 0000000..d24e33f --- /dev/null +++ b/sofic/automata/vpa/canonical.py @@ -0,0 +1,280 @@ +"""Canonical visibly pushdown automata from the well-matched summary algebra.""" + +from __future__ import annotations + +from collections import deque +from collections.abc import Hashable, Mapping, Sequence +from typing import Any + +from sofic.automata.vpa.base import VisiblyPushdownAutomaton +from sofic.automata.vpa.deterministic import DeterministicVisiblyPushdownAutomaton +from sofic.exceptions import NonWellMatchedLanguageError + + +class CanonicalVisiblyPushdownAutomaton(DeterministicVisiblyPushdownAutomaton): + """Canonical VPA built from the finite Myhill-Nerode summary algebra.""" + + summary_representatives: dict[Hashable, tuple[int | None, ...]] + + def __init__( + self, + *, + summary_representatives: Mapping[Hashable, tuple[int | None, ...]] | None = None, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + self.summary_representatives = dict(summary_representatives or {}) + + @classmethod + def from_vpa(cls, vpa: VisiblyPushdownAutomaton) -> CanonicalVisiblyPushdownAutomaton: + """Build the Myhill-Nerode canonical deterministic VPA of a well-matched language. + + States are classes of the summary algebra of well-matched factors, so + the form is canonical for well-matched languages. Raises + :class:`~sofic.exceptions.NonWellMatchedLanguageError` when ``vpa`` + accepts a word with a pending call or return; general VPLs have no + unique minimal deterministic VPA :cite:`AlurKumarMadhusudanViswanathan2005`. + With an empty call alphabet this is the minimal DFA. + """ + if vpa.has_unmatched_word(): + raise NonWellMatchedLanguageError( + "the canonical VPA is defined for well-matched languages; this one accepts a pending call or return" + ) + det = DeterministicVisiblyPushdownAutomaton.from_vpa(vpa) + return _SummaryAlgebra(det).to_canonical_vpa(cls) + + @classmethod + def minimize(cls, vpa: VisiblyPushdownAutomaton) -> CanonicalVisiblyPushdownAutomaton: + """Alias for :meth:`from_vpa`.""" + return cls.from_vpa(vpa) + + +class _SummaryAlgebra: + """Finite algebra of well-matched summaries with top-level and nested congruences. + + A summary maps each source state to the state reached along a well-matched + word (``None`` when the run dies). Top-level summaries are identified when + every well-matched continuation accepts both or neither; nested summaries + (inside a pending call) are identified when, for every enclosing context, + returning from them lands in the same class. Both partitions are refined + together until internal steps, calls, and returns are well defined on + classes, which makes the quotient canonical. + """ + + def __init__(self, vpa: DeterministicVisiblyPushdownAutomaton) -> None: + self.vpa = vpa + self.state_order = tuple(sorted(vpa.states(), key=repr)) + state_index = {state: index for index, state in enumerate(self.state_order)} + self.internal_summaries = { + symbol: _internal_summary(vpa, self.state_order, state_index, symbol) + for symbol in sorted(vpa.internal_alphabet, key=repr) + } + self.identity = tuple(range(len(self.state_order))) + self.summaries = sorted( + _close_summary_algebra(vpa, self.state_order, state_index, self.identity, self.internal_summaries), + key=repr, + ) + self.calls = sorted(vpa.call_alphabet, key=repr) + self.returns = sorted(vpa.return_alphabet, key=repr) + self._wrap_cache: dict[tuple, tuple] = {} + self.top, self.nested = self._refine() + + def wrap(self, inner: tuple, call: Any, ret: Any) -> tuple: + key = (inner, call, ret) + if key not in self._wrap_cache: + self._wrap_cache[key] = _wrap_summary(self.vpa, self.state_order, inner, call, ret) + return self._wrap_cache[key] + + def accepts(self, summary: tuple) -> bool: + return _summary_accepts(self.vpa, self.state_order, summary) + + def _refine(self) -> tuple[dict[tuple, int], dict[tuple, int]]: + top = {s: int(self.accepts(s)) for s in self.summaries} + nested = dict.fromkeys(self.summaries, 0) + internals = list(self.internal_summaries.values()) + pairs = [(c, r) for c in self.calls for r in self.returns] + while True: + top_signature = { + s: ( + top[s], + tuple(top[_compose_summary(s, a)] for a in internals), + tuple(top[_compose_summary(s, self.wrap(x, c, r))] for x in self.summaries for c, r in pairs), + ) + for s in self.summaries + } + nested_signature = { + s: ( + nested[s], + tuple(nested[_compose_summary(s, a)] for a in internals), + tuple(nested[_compose_summary(s, self.wrap(x, c, r))] for x in self.summaries for c, r in pairs), + tuple( + (top[_compose_summary(o, self.wrap(s, c, r))], nested[_compose_summary(o, self.wrap(s, c, r))]) + for o in self.summaries + for c, r in pairs + ), + ) + for s in self.summaries + } + new_top = _number_blocks(top_signature, first=top_signature[self.identity]) + new_nested = _number_blocks(nested_signature, first=nested_signature[self.identity]) + stable = len(set(new_top.values())) == len(set(top.values())) and len(set(new_nested.values())) == len( + set(nested.values()) + ) + top, nested = new_top, new_nested + if stable: + return top, nested + + def to_canonical_vpa(self, cls: type[CanonicalVisiblyPushdownAutomaton]) -> CanonicalVisiblyPushdownAutomaton: + representative: dict[Hashable, tuple] = {} + for s in self.summaries: + representative.setdefault(self.top[s], s) + representative.setdefault(("nested", self.nested[s]), s) + + def label(summary: tuple, nested: bool) -> Hashable: + return ("nested", self.nested[summary]) if nested else self.top[summary] + + initial = label(self.identity, False) + entry = label(self.identity, True) + reachable: set[Hashable] = set() + internals, calls, returns = set(), set(), set() + pushed: set[tuple[Hashable, Any]] = set() + frontier = [initial] + while frontier: + while frontier: + state = frontier.pop() + if state in reachable: + continue + reachable.add(state) + summary = representative[state] + for symbol, step in self.internal_summaries.items(): + target = label(_compose_summary(summary, step), isinstance(state, tuple)) + internals.add((state, symbol, target)) + frontier.append(target) + for call in self.calls: + calls.add((state, call, entry, (state, call))) + pushed.add((state, call)) + frontier.append(entry) + # Returns pop a pushed (outer, call) from a nested state; both sets + # grow together, so repeat until no new state appears. + for state in [s for s in reachable if isinstance(s, tuple)]: + for outer, call in list(pushed): + for ret in self.returns: + combined = _compose_summary(representative[outer], self.wrap(representative[state], call, ret)) + target = label(combined, isinstance(outer, tuple)) + returns.add((state, ret, (outer, call), target)) + if target not in reachable: + frontier.append(target) + + result = cls( + input_alphabet=self.vpa.input_alphabet, + call_alphabet=self.vpa.call_alphabet, + return_alphabet=self.vpa.return_alphabet, + internal_alphabet=self.vpa.internal_alphabet, + stack_alphabet=frozenset(pushed), + bottom_stack_symbol=None, + initial_state=initial, + accepting_states=frozenset( + s for s in reachable if not isinstance(s, tuple) and self.accepts(representative[s]) + ), + summary_representatives={s: representative[s] for s in reachable}, + ) + for state in sorted(reachable, key=repr): + result.graph.add_state(state) + for source, symbol, target in sorted(internals, key=repr): + result.add_internal_transition(source, target, symbol) + for source, symbol, target, push in sorted(calls, key=repr): + result.add_call_transition(source, target, symbol, push) + for source, symbol, guard, target in sorted(returns, key=repr): + result.add_return_transition(source, target, symbol, guard) + result.validate() + return result + + +def _number_blocks(signatures: Mapping[tuple, Any], *, first: Any) -> dict[tuple, int]: + ordered = sorted(set(signatures.values()), key=lambda sig: (sig != first, repr(sig))) + index = {signature: position for position, signature in enumerate(ordered)} + return {summary: index[signature] for summary, signature in signatures.items()} + + +def _internal_summary( + vpa: DeterministicVisiblyPushdownAutomaton, + state_order: Sequence[Hashable], + state_index: Mapping[Hashable, int], + symbol: Any, +) -> tuple[int | None, ...]: + transitions = vpa.internal_transition_map() + summary: list[int | None] = [] + for state in state_order: + target = transitions.get((state, symbol)) + summary.append(None if target is None else state_index[target]) + return tuple(summary) + + +def _close_summary_algebra( + vpa: DeterministicVisiblyPushdownAutomaton, + state_order: Sequence[Hashable], + state_index: Mapping[Hashable, int], + identity: tuple[int | None, ...], + internal_summaries: Mapping[Any, tuple[int | None, ...]], +) -> set[tuple[int | None, ...]]: + summaries = {identity, *internal_summaries.values()} + queue: deque[tuple[int | None, ...]] = deque(sorted(summaries, key=repr)) + while queue: + summary = queue.popleft() + current = list(summaries) + candidates: list[tuple[int | None, ...]] = [] + for other in current: + candidates.append(_compose_summary(summary, other)) + candidates.append(_compose_summary(other, summary)) + for call_symbol in sorted(vpa.call_alphabet, key=repr): + for return_symbol in sorted(vpa.return_alphabet, key=repr): + candidates.append(_wrap_summary(vpa, state_order, summary, call_symbol, return_symbol)) + for candidate in candidates: + if candidate not in summaries: + summaries.add(candidate) + queue.append(candidate) + return summaries + + +def _compose_summary( + first: tuple[int | None, ...], + second: tuple[int | None, ...], +) -> tuple[int | None, ...]: + return tuple(None if state is None else second[state] for state in first) + + +def _wrap_summary( + vpa: DeterministicVisiblyPushdownAutomaton, + state_order: Sequence[Hashable], + inner: tuple[int | None, ...], + call_symbol: Any, + return_symbol: Any, +) -> tuple[int | None, ...]: + state_index = {state: index for index, state in enumerate(state_order)} + call_map = vpa.call_transition_map() + result: list[int | None] = [] + for state in state_order: + call = call_map.get((state, call_symbol)) + if call is None: + result.append(None) + continue + call_target, stack_symbol = call + inner_target_index = inner[state_index[call_target]] + if inner_target_index is None: + result.append(None) + continue + return_target = vpa.return_successor(state_order[inner_target_index], return_symbol, stack_symbol) + result.append(None if return_target is None else state_index[return_target]) + return tuple(result) + + +def _summary_accepts( + vpa: DeterministicVisiblyPushdownAutomaton, + state_order: Sequence[Hashable], + summary: tuple[int | None, ...], +) -> bool: + if vpa.initial_state is None: + return False + initial_index = {state: index for index, state in enumerate(state_order)}[vpa.initial_state] + target = summary[initial_index] + return target is not None and state_order[target] in vpa.accepting_states diff --git a/sofic/automata/vpa/deterministic.py b/sofic/automata/vpa/deterministic.py new file mode 100644 index 0000000..c1916f9 --- /dev/null +++ b/sofic/automata/vpa/deterministic.py @@ -0,0 +1,166 @@ +"""Deterministic visibly pushdown automata.""" + +from __future__ import annotations + +from collections.abc import Hashable +from typing import Any + +from sofic.automata.vpa import operations as vc +from sofic.automata.vpa.base import _MISSING, VisiblyPushdownAutomaton +from sofic.exceptions import NonDeterministicError +from sofic.graph import ( + ATTR_KIND, + ATTR_STACK_SYMBOL, + ATTR_SYMBOL, + KIND_CALL, + KIND_INTERNAL, + KIND_RETURN, +) + + +class DeterministicVisiblyPushdownAutomaton(VisiblyPushdownAutomaton): + """VPA with at most one enabled transition for each visible configuration.""" + + def validate(self) -> None: + super().validate() + if self.initial_state is None: + raise NonDeterministicError("deterministic VPA requires an initial state") + self._check_determinism() + + def _check_determinism(self) -> None: + call_keys: set[tuple[Hashable, Any]] = set() + internal_keys: set[tuple[Hashable, Any]] = set() + return_keys: set[tuple[Hashable, Any, Any]] = set() + wildcard_returns: set[tuple[Hashable, Any]] = set() + + for transition in self.transitions(): + kind = transition.data.get(ATTR_KIND) + symbol = transition.data.get(ATTR_SYMBOL) + if symbol is None: + continue + if kind == KIND_CALL: + key = (transition.source, symbol) + if key in call_keys: + raise NonDeterministicError(f"non-deterministic call transition on {key}") + call_keys.add(key) + elif kind == KIND_INTERNAL: + key = (transition.source, symbol) + if key in internal_keys: + raise NonDeterministicError(f"non-deterministic internal transition on {key}") + internal_keys.add(key) + elif kind == KIND_RETURN: + stack_symbol = transition.data.get(ATTR_STACK_SYMBOL) + wildcard_key = (transition.source, symbol) + if stack_symbol is None: + if wildcard_key in wildcard_returns: + raise NonDeterministicError(f"duplicate wildcard return transition on {wildcard_key}") + if any(source == transition.source and ret == symbol for source, ret, _stack in return_keys): + raise NonDeterministicError( + f"wildcard return overlaps guarded return on {wildcard_key}", + ) + wildcard_returns.add(wildcard_key) + else: + key = (transition.source, symbol, stack_symbol) + if wildcard_key in wildcard_returns: + raise NonDeterministicError( + f"guarded return overlaps wildcard return on {wildcard_key}", + ) + if key in return_keys: + raise NonDeterministicError(f"duplicate guarded return transition on {key}") + return_keys.add(key) + + def add_call_transition( + self, + source: Hashable, + target: Hashable, + symbol: Any, + stack_symbol: Any, + **attrs: Any, + ) -> int: + for transition in self.graph.out_transitions(source): + if transition.data.get(ATTR_KIND) == KIND_CALL and transition.data.get(ATTR_SYMBOL) == symbol: + raise NonDeterministicError(f"non-deterministic call transition on {(source, symbol)}") + return super().add_call_transition(source, target, symbol, stack_symbol, **attrs) + + def add_return_transition( + self, + source: Hashable, + target: Hashable, + symbol: Any, + stack_symbol: Any | None = None, + **attrs: Any, + ) -> int: + for transition in self.graph.out_transitions(source): + if transition.data.get(ATTR_KIND) != KIND_RETURN or transition.data.get(ATTR_SYMBOL) != symbol: + continue + existing_stack = transition.data.get(ATTR_STACK_SYMBOL) + if existing_stack is None or stack_symbol is None or existing_stack == stack_symbol: + raise NonDeterministicError(f"non-deterministic return transition on {(source, symbol)}") + return super().add_return_transition(source, target, symbol, stack_symbol, **attrs) + + def add_internal_transition(self, source: Hashable, target: Hashable, symbol: Any, **attrs: Any) -> int: + for transition in self.graph.out_transitions(source): + if transition.data.get(ATTR_KIND) == KIND_INTERNAL and transition.data.get(ATTR_SYMBOL) == symbol: + raise NonDeterministicError(f"non-deterministic internal transition on {(source, symbol)}") + return super().add_internal_transition(source, target, symbol, **attrs) + + def call_successor(self, state: Hashable, symbol: Any) -> tuple[Hashable, Any] | None: + """Return ``(target, pushed_stack_symbol)`` for a deterministic call.""" + return self.call_transition_map().get((state, symbol)) + + def internal_successor(self, state: Hashable, symbol: Any) -> Hashable | None: + """Return the deterministic internal successor, if present.""" + return self.internal_transition_map().get((state, symbol)) + + def return_successor(self, state: Hashable, symbol: Any, stack_symbol: Any) -> Hashable | None: + """Return the deterministic return successor for ``stack_symbol``, if present.""" + transitions = self.return_transition_map() + explicit = transitions.get((state, symbol, stack_symbol), _MISSING) + if explicit is not _MISSING: + return explicit + wildcard = transitions.get((state, symbol, None), _MISSING) + if wildcard is not _MISSING: + return wildcard + return None + + @classmethod + def from_vpa(cls, vpa: VisiblyPushdownAutomaton) -> DeterministicVisiblyPushdownAutomaton: + """Return ``vpa`` as a deterministic VPA, determinizing it when needed. + + An already deterministic ``vpa`` is copied with its states unchanged; + otherwise the summary construction of :cite:`AlurMadhusudan2009` is + applied (see :func:`~sofic.automata.vpa.operations.determinize`). + """ + if vpa.initial_state is None or not _is_deterministic(vpa): + return vc.determinize_vpa(vpa) + result = cls( + input_alphabet=vpa.input_alphabet, + call_alphabet=vpa.call_alphabet, + return_alphabet=vpa.return_alphabet, + internal_alphabet=vpa.internal_alphabet, + stack_alphabet=vpa.stack_alphabet, + bottom_stack_symbol=vpa.bottom_stack_symbol, + initial_state=vpa.initial_state, + accepting_states=vpa.accepting_states, + graph=vpa.graph.copy(), + ) + result.validate() + return result + + +def _is_deterministic(vpa: VisiblyPushdownAutomaton) -> bool: + probe = DeterministicVisiblyPushdownAutomaton( + call_alphabet=vpa.call_alphabet, + return_alphabet=vpa.return_alphabet, + internal_alphabet=vpa.internal_alphabet, + stack_alphabet=vpa.stack_alphabet, + bottom_stack_symbol=vpa.bottom_stack_symbol, + initial_state=vpa.initial_state, + accepting_states=vpa.accepting_states, + graph=vpa.graph, + ) + try: + probe._check_determinism() + except NonDeterministicError: + return False + return True diff --git a/sofic/automata/vpa/modular.py b/sofic/automata/vpa/modular.py new file mode 100644 index 0000000..ec6759e --- /dev/null +++ b/sofic/automata/vpa/modular.py @@ -0,0 +1,686 @@ +"""Modular (call-driven, single- and multiple-entry) visibly pushdown automata.""" + +from __future__ import annotations + +from collections.abc import Hashable, Iterable, Mapping, Sequence +from typing import Any + +from sofic.automata.vpa import operations as vc +from sofic.automata.vpa.base import _MISSING, VisiblyPushdownAutomaton +from sofic.automata.vpa.deterministic import DeterministicVisiblyPushdownAutomaton +from sofic.exceptions import NonDeterministicError +from sofic.graph import ( + ATTR_KIND, + ATTR_STACK_SYMBOL, + ATTR_SYMBOL, + KIND_CALL, + KIND_INTERNAL, + KIND_RETURN, +) + + +class CallDrivenAutomaton(DeterministicVisiblyPushdownAutomaton): + """Deterministic modular VPA whose call target depends only on the call symbol.""" + + modules: dict[Hashable, frozenset[Hashable]] + base_module: Hashable + call_partition: dict[Any, Hashable] + call_entries: dict[Any, Hashable] + + def __init__( + self, + *, + modules: Mapping[Hashable, Iterable[Hashable]] | None = None, + base_module: Hashable = 0, + call_partition: Mapping[Any, Hashable] | None = None, + call_entries: Mapping[Any, Hashable] | None = None, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + self.modules = _normalize_modules(modules) + self.base_module = base_module + self.call_partition = dict(call_partition or {}) + self.call_entries = dict(call_entries or {}) + + def validate(self) -> None: + super().validate() + state_modules = self._validate_modules() + self._validate_call_partition() + self._validate_internal_transitions_stay_in_module(state_modules) + self._validate_call_driven_transitions() + + def call_entry_map(self) -> dict[Any, Hashable]: + """Return configured or inferred entries for each call symbol with transitions.""" + entries = dict(self.call_entries) + for transition in self.transitions(): + if transition.data.get(ATTR_KIND) != KIND_CALL: + continue + symbol = transition.data.get(ATTR_SYMBOL) + if symbol is None: + continue + existing = entries.get(symbol, _MISSING) + if existing is _MISSING: + entries[symbol] = transition.target + elif existing != transition.target: + raise NonDeterministicError(f"call target for {symbol!r} depends on source state") + return entries + + @classmethod + def minimize( + cls, + vpa: VisiblyPushdownAutomaton, + *, + modules: Mapping[Hashable, Iterable[Hashable]] | None = None, + call_partition: Mapping[Any, Hashable] | None = None, + base_module: Hashable | None = None, + call_entries: Mapping[Any, Hashable] | None = None, + ) -> CallDrivenAutomaton: + """Return the module-aware deterministic quotient as a CDA.""" + return _minimize_modular_vpa( + cls, + vpa, + modules=modules, + call_partition=call_partition, + base_module=base_module, + call_entries=call_entries, + entry_states=None, + form="cda", + ) + + def _validate_modules(self) -> dict[Hashable, Hashable]: + self._require(bool(self.modules), "modular VPA requires modules") + self._require(self.base_module in self.modules, "base_module must be present in modules") + all_states = set(self.states()) + seen: dict[Hashable, Hashable] = {} + for module, states in self.modules.items(): + self._require(bool(states), f"module {module!r} must contain at least one state") + for state in states: + self._require(state in all_states, f"module {module!r} contains unknown state {state!r}") + self._require(state not in seen, f"state {state!r} appears in multiple modules") + seen[state] = module + self._require(set(seen) == all_states, "modules must cover exactly the VPA states") + if self.initial_state is not None: + self._require( + self.initial_state in self.modules[self.base_module], + "initial_state must lie in the base module", + ) + return seen + + def _validate_call_partition(self) -> None: + missing = self.call_alphabet - set(self.call_partition) + extra = set(self.call_partition) - self.call_alphabet + self._require(not missing, f"call_partition missing calls {sorted(missing, key=repr)!r}") + self._require(not extra, f"call_partition contains non-call symbols {sorted(extra, key=repr)!r}") + for symbol, module in self.call_partition.items(): + self._require(module in self.modules, f"call {symbol!r} targets unknown module {module!r}") + + def _validate_internal_transitions_stay_in_module(self, state_modules: Mapping[Hashable, Hashable]) -> None: + for transition in self.transitions(): + if transition.data.get(ATTR_KIND) == KIND_INTERNAL: + self._require( + state_modules[transition.source] == state_modules[transition.target], + "internal transitions must stay inside one module", + ) + + def _validate_call_driven_transitions(self) -> None: + entries = self.call_entry_map() + for symbol, target in entries.items(): + module = self.call_partition[symbol] + self._require(target in self.modules[module], f"call {symbol!r} entry is not in module {module!r}") + for transition in self.transitions(): + if transition.data.get(ATTR_KIND) != KIND_CALL: + continue + symbol = transition.data.get(ATTR_SYMBOL) + if symbol is None: + continue + self._require( + transition.target == entries[symbol], + f"call target for {symbol!r} must be independent of source state", + ) + + +class MultipleEntryVisiblyPushdownAutomaton(CallDrivenAutomaton): + """Modular VPA with multiple module entries and source-determined call pushes.""" + + entry_states: dict[Hashable, frozenset[Hashable]] + + def __init__( + self, + *, + entry_states: Mapping[Hashable, Iterable[Hashable] | Hashable] | None = None, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + self.entry_states = _normalize_multi_entries(entry_states) + + def validate(self) -> None: + super().validate() + self._validate_entry_states() + self._validate_call_targets_are_entries() + self._validate_source_determined_pushes() + + @classmethod + def minimize( + cls, + vpa: VisiblyPushdownAutomaton, + *, + modules: Mapping[Hashable, Iterable[Hashable]] | None = None, + call_partition: Mapping[Any, Hashable] | None = None, + base_module: Hashable | None = None, + entry_states: Mapping[Hashable, Iterable[Hashable] | Hashable] | None = None, + call_entries: Mapping[Any, Hashable] | None = None, + ) -> MultipleEntryVisiblyPushdownAutomaton: + """Return the module-aware deterministic quotient as an MEVPA.""" + return _minimize_modular_vpa( + cls, + vpa, + modules=modules, + call_partition=call_partition, + base_module=base_module, + call_entries=call_entries, + entry_states=entry_states, + form="mevpa", + ) + + def _validate_entry_states(self) -> None: + self._require(bool(self.entry_states), "MEVPA requires entry_states") + for module, states in self.entry_states.items(): + self._require(module in self.modules, f"entry_states has unknown module {module!r}") + self._require(bool(states), f"module {module!r} must have at least one entry") + for state in states: + self._require(state in self.modules[module], f"entry {state!r} is not in module {module!r}") + + def _validate_call_targets_are_entries(self) -> None: + entries = self.call_entry_map() + for symbol, target in entries.items(): + module = self.call_partition[symbol] + self._require( + target in self.entry_states.get(module, frozenset()), + f"call {symbol!r} must enter one of module {module!r}'s entries", + ) + + def _validate_source_determined_pushes(self) -> None: + pushed_by_source: dict[Hashable, Any] = {} + for transition in self.transitions(): + if transition.data.get(ATTR_KIND) != KIND_CALL: + continue + pushed = transition.data.get(ATTR_STACK_SYMBOL) + existing = pushed_by_source.get(transition.source, _MISSING) + if existing is _MISSING: + pushed_by_source[transition.source] = pushed + else: + self._require(existing == pushed, "MEVPA call push must depend only on the source state") + + +class SingleEntryVisiblyPushdownAutomaton(CallDrivenAutomaton): + """Modular VPA with one distinguished entry per non-base module.""" + + entry_states: dict[Hashable, Hashable] + + def __init__( + self, + *, + entry_states: Mapping[Hashable, Hashable] | None = None, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + self.entry_states = dict(entry_states or {}) + + def validate(self) -> None: + super().validate() + self._validate_single_entries() + self._validate_single_entry_calls() + + @classmethod + def minimize( + cls, + vpa: VisiblyPushdownAutomaton, + *, + call_partition: Mapping[Any, Hashable] | None = None, + modules: Mapping[Hashable, Iterable[Hashable]] | None = None, + base_module: Hashable | None = None, + entry_states: Mapping[Hashable, Hashable] | None = None, + call_entries: Mapping[Any, Hashable] | None = None, + ) -> SingleEntryVisiblyPushdownAutomaton: + """Return the module-aware deterministic quotient as an SEVPA. + + A fixed call partition and module structure are required. General VPA + minimization is intentionally not attempted here. + """ + return _minimize_modular_vpa( + cls, + vpa, + modules=modules, + call_partition=call_partition, + base_module=base_module, + call_entries=call_entries, + entry_states=entry_states, + form="sevpa", + ) + + def _validate_single_entries(self) -> None: + missing = set(self.modules) - {self.base_module} - set(self.entry_states) + self._require(not missing, f"SEVPA missing entries for modules {sorted(missing, key=repr)!r}") + for module, state in self.entry_states.items(): + self._require(module in self.modules, f"entry_states has unknown module {module!r}") + self._require(state in self.modules[module], f"entry {state!r} is not in module {module!r}") + + def _validate_single_entry_calls(self) -> None: + for transition in self.transitions(): + if transition.data.get(ATTR_KIND) != KIND_CALL: + continue + symbol = transition.data.get(ATTR_SYMBOL) + if symbol is None: + continue + module = self.call_partition[symbol] + expected_entry = self.entry_states.get(module) + self._require( + transition.target == expected_entry, + f"SEVPA call {symbol!r} must enter module {module!r}'s single entry", + ) + self._require( + transition.data.get(ATTR_STACK_SYMBOL) == (transition.source, symbol), + "SEVPA call stack symbols must be (caller_state, call_symbol)", + ) + + +def _normalize_modules(modules: Mapping[Hashable, Iterable[Hashable]] | None) -> dict[Hashable, frozenset[Hashable]]: + if modules is None: + return {} + return {module: frozenset(states) for module, states in modules.items()} + + +def _normalize_multi_entries( + entry_states: Mapping[Hashable, Iterable[Hashable] | Hashable] | None, +) -> dict[Hashable, frozenset[Hashable]]: + if entry_states is None: + return {} + result: dict[Hashable, frozenset[Hashable]] = {} + for module, states in entry_states.items(): + if isinstance(states, frozenset | set | list): + result[module] = frozenset(states) + else: + result[module] = frozenset({states}) + return result + + +def _metadata_or_argument(vpa: VisiblyPushdownAutomaton, name: str, value: Any, default: Any) -> Any: + if value is not None: + return value + return getattr(vpa, name, default) + + +def _deterministic_view(vpa: VisiblyPushdownAutomaton) -> DeterministicVisiblyPushdownAutomaton: + return DeterministicVisiblyPushdownAutomaton.from_vpa(vpa) + + +def _minimize_modular_vpa( + target_cls: type[CallDrivenAutomaton], + vpa: VisiblyPushdownAutomaton, + *, + modules: Mapping[Hashable, Iterable[Hashable]] | None, + call_partition: Mapping[Any, Hashable] | None, + base_module: Hashable | None, + call_entries: Mapping[Any, Hashable] | None, + entry_states: Mapping[Hashable, Any] | None, + form: str, +) -> Any: + modules = _metadata_or_argument(vpa, "modules", modules, None) + call_partition = _metadata_or_argument(vpa, "call_partition", call_partition, None) + base_module = _metadata_or_argument(vpa, "base_module", base_module, 0) + call_entries = _metadata_or_argument(vpa, "call_entries", call_entries, None) + entry_states = _metadata_or_argument(vpa, "entry_states", entry_states, None) + + if modules is None: + convert = vc.to_multiple_entry if form == "mevpa" else vc.to_single_entry + converted = convert(vpa, call_partition) + return _minimize_modular_vpa( + target_cls, + converted, + modules=converted.modules, + call_partition=converted.call_partition, + base_module=converted.base_module, + call_entries=converted.call_entries, + entry_states=getattr(converted, "entry_states", None) if form != "cda" else None, + form=form, + ) + if call_partition is None: + raise ValueError("a call_partition is required when modules are given") + + det = _deterministic_view(vpa) + modules = _normalize_modules(modules) + call_partition = dict(call_partition) + call_entries = dict(call_entries or _infer_call_entries(det, call_partition)) + entry_states = _infer_entry_states(form, modules, base_module, call_partition, call_entries, entry_states, det) + + source = _construct_modular_view( + target_cls, + det, + modules=modules, + base_module=base_module, + call_partition=call_partition, + call_entries=call_entries, + entry_states=entry_states, + ) + source.validate() + + partition = _refine_modular_partition(source, modules) + return _quotient_modular_vpa( + target_cls, + source, + partition, + modules=modules, + base_module=base_module, + call_partition=call_partition, + call_entries=call_entries, + entry_states=entry_states, + form=form, + ) + + +def _construct_modular_view( + target_cls: type[CallDrivenAutomaton], + det: DeterministicVisiblyPushdownAutomaton, + *, + modules: Mapping[Hashable, frozenset[Hashable]], + base_module: Hashable, + call_partition: Mapping[Any, Hashable], + call_entries: Mapping[Any, Hashable], + entry_states: Mapping[Hashable, Any], +) -> CallDrivenAutomaton: + kwargs = { + "input_alphabet": det.input_alphabet, + "call_alphabet": det.call_alphabet, + "return_alphabet": det.return_alphabet, + "internal_alphabet": det.internal_alphabet, + "stack_alphabet": det.stack_alphabet, + "bottom_stack_symbol": det.bottom_stack_symbol, + "initial_state": det.initial_state, + "accepting_states": det.accepting_states, + "graph": det.graph.copy(), + "modules": modules, + "base_module": base_module, + "call_partition": call_partition, + "call_entries": call_entries, + } + if issubclass(target_cls, (SingleEntryVisiblyPushdownAutomaton, MultipleEntryVisiblyPushdownAutomaton)): + kwargs["entry_states"] = entry_states + return target_cls(**kwargs) + + +def _infer_call_entries( + det: DeterministicVisiblyPushdownAutomaton, + call_partition: Mapping[Any, Hashable], +) -> dict[Any, Hashable]: + entries: dict[Any, Hashable] = {} + for (_source, symbol), (target, _stack) in det.call_transition_map().items(): + if symbol not in call_partition: + raise ValueError(f"call_partition missing call symbol {symbol!r}") + existing = entries.get(symbol, _MISSING) + if existing is _MISSING: + entries[symbol] = target + elif existing != target: + raise ValueError( + f"call target for {symbol!r} depends on the source state; omit modules to convert the VPA first" + ) + return entries + + +def _infer_entry_states( + form: str, + modules: Mapping[Hashable, frozenset[Hashable]], + base_module: Hashable, + call_partition: Mapping[Any, Hashable], + call_entries: Mapping[Any, Hashable], + entry_states: Mapping[Hashable, Any] | None, + det: DeterministicVisiblyPushdownAutomaton, +) -> dict[Hashable, Any]: + if entry_states is not None: + if form == "mevpa": + return _normalize_multi_entries(entry_states) + return dict(entry_states) + + by_module: dict[Hashable, set[Hashable]] = {module: set() for module in modules} + if det.initial_state is not None and base_module in by_module: + by_module[base_module].add(det.initial_state) + for symbol, entry in call_entries.items(): + by_module[call_partition[symbol]].add(entry) + + if form == "mevpa": + return {module: frozenset(states) for module, states in by_module.items() if states} + if form == "sevpa": + result: dict[Hashable, Hashable] = {} + for module, states in by_module.items(): + if module == base_module: + continue + if len(states) != 1: + raise ValueError( + "SEVPA minimization requires one entry per non-base module; omit modules to convert the VPA first" + ) + result[module] = next(iter(states)) + return result + return {} + + +def _refine_modular_partition( + vpa: CallDrivenAutomaton, + modules: Mapping[Hashable, frozenset[Hashable]], +) -> list[frozenset[Hashable]]: + partition: list[frozenset[Hashable]] = [] + for _module, states in sorted(modules.items(), key=lambda item: repr(item[0])): + accepting = frozenset(states & vpa.accepting_states) + rejecting = frozenset(states - vpa.accepting_states) + if accepting: + partition.append(accepting) + if rejecting: + partition.append(rejecting) + + changed = True + while changed: + changed = False + block_of = _block_map(partition) + context_groups = _stack_context_groups(vpa, block_of) + new_partition: list[frozenset[Hashable]] = [] + for block in partition: + pieces: dict[tuple[Any, ...], set[Hashable]] = {} + for state in block: + signature = _modular_state_signature(vpa, state, block_of, context_groups) + pieces.setdefault(signature, set()).add(state) + if len(pieces) > 1: + changed = True + new_partition.extend(frozenset(piece) for piece in pieces.values()) + partition = new_partition + return partition + + +def _block_map(partition: Sequence[frozenset[Hashable]]) -> dict[Hashable, frozenset[Hashable]]: + return {state: block for block in partition for state in block} + + +def _stack_context_groups( + vpa: DeterministicVisiblyPushdownAutomaton, + block_of: Mapping[Hashable, Hashable], +) -> list[tuple[Any, tuple[Any, ...]]]: + stack_symbols = {stack for _key, (_target, stack) in vpa.call_transition_map().items()} + if vpa.bottom_stack_symbol is not None: + stack_symbols.add(vpa.bottom_stack_symbol) + grouped: dict[Any, set[Any]] = {} + for stack_symbol in stack_symbols: + canonical = _canonical_stack_symbol(stack_symbol, block_of) + grouped.setdefault(canonical, set()).add(stack_symbol) + return [(canonical, tuple(sorted(actuals, key=repr))) for canonical, actuals in sorted(grouped.items(), key=repr)] + + +def _modular_state_signature( + vpa: CallDrivenAutomaton, + state: Hashable, + block_of: Mapping[Hashable, Hashable], + context_groups: Sequence[tuple[Any, tuple[Any, ...]]], +) -> tuple[Any, ...]: + state_modules = {state: module for module, states in vpa.modules.items() for state in states} + call_map = vpa.call_transition_map() + internal_map = vpa.internal_transition_map() + return_map = vpa.return_transition_map() + + internal = tuple( + ( + symbol, + None if (target := internal_map.get((state, symbol))) is None else block_of[target], + ) + for symbol in sorted(vpa.internal_alphabet, key=repr) + ) + calls = tuple( + ( + symbol, + None + if (call := call_map.get((state, symbol))) is None + else (block_of[call[0]], _canonical_stack_symbol(call[1], block_of)), + ) + for symbol in sorted(vpa.call_alphabet, key=repr) + ) + returns = [] + for symbol in sorted(vpa.return_alphabet, key=repr): + for canonical_stack, actual_stacks in context_groups: + targets = [] + for stack_symbol in actual_stacks: + target = return_map.get((state, symbol, stack_symbol), return_map.get((state, symbol, None))) + targets.append(None if target is None else block_of[target]) + returns.append((symbol, canonical_stack, tuple(sorted(set(targets), key=repr)))) + + return ( + state_modules[state], + state in vpa.accepting_states, + internal, + calls, + tuple(returns), + ) + + +def _quotient_modular_vpa( + target_cls: type[CallDrivenAutomaton], + source: CallDrivenAutomaton, + partition: Sequence[frozenset[Hashable]], + *, + modules: Mapping[Hashable, frozenset[Hashable]], + base_module: Hashable, + call_partition: Mapping[Any, Hashable], + call_entries: Mapping[Any, Hashable], + entry_states: Mapping[Hashable, Any], + form: str, +) -> Any: + block_of = _block_map(partition) + blocks = list(partition) + stack_alphabet: set[Any] = set() + if source.bottom_stack_symbol is not None: + stack_alphabet.add(source.bottom_stack_symbol) + transitions: set[tuple[Hashable, Hashable, Any, Any, Any | None]] = set() + stack_symbol_rewrite: dict[Any, set[Any]] = {} + + for transition in source.transitions(): + if transition.data.get(ATTR_KIND) != KIND_CALL: + continue + symbol = transition.data.get(ATTR_SYMBOL) + stack_symbol = transition.data.get(ATTR_STACK_SYMBOL) + quotient_stack_symbol = _quotient_call_stack_symbol( + form, + block_of[transition.source], + symbol, + stack_symbol, + block_of, + ) + stack_symbol_rewrite.setdefault(stack_symbol, set()).add(quotient_stack_symbol) + + for block in blocks: + representative = min(block, key=repr) + source_block = block_of[representative] + for transition in source.graph.out_transitions(representative): + kind = transition.data.get(ATTR_KIND) + symbol = transition.data.get(ATTR_SYMBOL) + target_block = block_of[transition.target] + stack_symbol = transition.data.get(ATTR_STACK_SYMBOL) + if kind == KIND_CALL: + stack_symbol = _quotient_call_stack_symbol(form, source_block, symbol, stack_symbol, block_of) + stack_alphabet.add(stack_symbol) + elif kind == KIND_RETURN and stack_symbol is not None: + rewritten = stack_symbol_rewrite.get(stack_symbol, {_canonical_stack_symbol(stack_symbol, block_of)}) + for rewritten_stack_symbol in rewritten: + stack_alphabet.add(rewritten_stack_symbol) + transitions.add((source_block, target_block, kind, symbol, rewritten_stack_symbol)) + continue + transitions.add((source_block, target_block, kind, symbol, stack_symbol)) + + quotient_modules = { + module: frozenset(block for block in blocks if block & states) + for module, states in modules.items() + if any(block & states for block in blocks) + } + quotient_call_entries = {symbol: block_of[state] for symbol, state in call_entries.items() if state in block_of} + quotient_entry_states = _quotient_entry_states(form, entry_states, block_of) + + kwargs = { + "input_alphabet": source.input_alphabet, + "call_alphabet": source.call_alphabet, + "return_alphabet": source.return_alphabet, + "internal_alphabet": source.internal_alphabet, + "stack_alphabet": frozenset(stack_alphabet), + "bottom_stack_symbol": source.bottom_stack_symbol, + "initial_state": block_of[source.initial_state], + "accepting_states": frozenset(block for block in blocks if block & source.accepting_states), + "modules": quotient_modules, + "base_module": base_module, + "call_partition": dict(call_partition), + "call_entries": quotient_call_entries, + } + if issubclass(target_cls, (SingleEntryVisiblyPushdownAutomaton, MultipleEntryVisiblyPushdownAutomaton)): + kwargs["entry_states"] = quotient_entry_states + + result = target_cls(**kwargs) + for block in blocks: + result.graph.add_state(block) + for source_block, target_block, kind, symbol, stack_symbol in sorted(transitions, key=repr): + if kind == KIND_CALL: + result.add_call_transition(source_block, target_block, symbol, stack_symbol) + elif kind == KIND_RETURN: + result.add_return_transition(source_block, target_block, symbol, stack_symbol) + elif kind == KIND_INTERNAL: + result.add_internal_transition(source_block, target_block, symbol) + result.validate() + return result + + +def _quotient_call_stack_symbol( + form: str, + source_block: Hashable, + symbol: Any, + stack_symbol: Any, + block_of: Mapping[Hashable, Hashable], +) -> Any: + if form == "sevpa": + return (source_block, symbol) + if form == "mevpa": + return source_block + return _canonical_stack_symbol(stack_symbol, block_of) + + +def _quotient_entry_states( + form: str, + entry_states: Mapping[Hashable, Any], + block_of: Mapping[Hashable, Hashable], +) -> dict[Hashable, Any]: + if form == "mevpa": + normalized = _normalize_multi_entries(entry_states) + return { + module: frozenset(block_of[state] for state in states if state in block_of) + for module, states in normalized.items() + } + if form == "sevpa": + return {module: block_of[state] for module, state in entry_states.items() if state in block_of} + return {} + + +def _canonical_stack_symbol(stack_symbol: Any, block_of: Mapping[Hashable, Hashable]) -> Any: + if stack_symbol in block_of: + return block_of[stack_symbol] + if isinstance(stack_symbol, tuple): + return tuple(_canonical_stack_symbol(part, block_of) for part in stack_symbol) + return stack_symbol diff --git a/sofic/automata/vpa_constructions.py b/sofic/automata/vpa/operations.py similarity index 98% rename from sofic/automata/vpa_constructions.py rename to sofic/automata/vpa/operations.py index 494a39d..1a46dae 100644 --- a/sofic/automata/vpa_constructions.py +++ b/sofic/automata/vpa/operations.py @@ -36,12 +36,9 @@ from sofic.graph import ATTR_KIND, ATTR_STACK_SYMBOL, ATTR_SYMBOL, KIND_CALL, KIND_INTERNAL, KIND_RETURN if TYPE_CHECKING: - from sofic.automata.vpa import ( - DeterministicVisiblyPushdownAutomaton, - MultipleEntryVisiblyPushdownAutomaton, - SingleEntryVisiblyPushdownAutomaton, - VisiblyPushdownAutomaton, - ) + from sofic.automata.vpa.base import VisiblyPushdownAutomaton + from sofic.automata.vpa.deterministic import DeterministicVisiblyPushdownAutomaton + from sofic.automata.vpa.modular import MultipleEntryVisiblyPushdownAutomaton, SingleEntryVisiblyPushdownAutomaton #: Bottom-of-stack symbol used by constructed VPAs. BOTTOM = "⊥" @@ -178,7 +175,7 @@ def _forward_trim(machine: NormalVPA) -> NormalVPA: def denormalize(machine: NormalVPA, cls: type | None = None, **extra: Any) -> Any: """Build a VPA of class ``cls`` from a normalized form, relabeling states and stack symbols as integers.""" - from sofic.automata.vpa import VisiblyPushdownAutomaton + from sofic.automata.vpa.base import VisiblyPushdownAutomaton cls = cls or VisiblyPushdownAutomaton machine = _forward_trim(machine) @@ -553,7 +550,7 @@ def _modular_conversion( *, multiple_entry: bool, ) -> Any: - from sofic.automata.vpa import MultipleEntryVisiblyPushdownAutomaton, SingleEntryVisiblyPushdownAutomaton + from sofic.automata.vpa.modular import MultipleEntryVisiblyPushdownAutomaton, SingleEntryVisiblyPushdownAutomaton machine = normalize(vpa) if has_unmatched_word(machine): @@ -713,6 +710,6 @@ def to_multiple_entry( def determinize_vpa(vpa: VisiblyPushdownAutomaton) -> DeterministicVisiblyPushdownAutomaton: - from sofic.automata.vpa import DeterministicVisiblyPushdownAutomaton + from sofic.automata.vpa.deterministic import DeterministicVisiblyPushdownAutomaton return denormalize(determinize(normalize(vpa)), DeterministicVisiblyPushdownAutomaton) diff --git a/sofic/automata/vpa_simulation.py b/sofic/automata/vpa/simulation.py similarity index 87% rename from sofic/automata/vpa_simulation.py rename to sofic/automata/vpa/simulation.py index f40e1cb..b0bc2d9 100644 --- a/sofic/automata/vpa_simulation.py +++ b/sofic/automata/vpa/simulation.py @@ -6,8 +6,8 @@ from typing import Any from sofic.automata._config_simulation import simulate_configs -from sofic.automata.vpa import VisiblyPushdownAutomaton -from sofic.automata.vpa_constructions import BOTTOM, normalize +from sofic.automata.vpa.base import VisiblyPushdownAutomaton +from sofic.automata.vpa.operations import BOTTOM, normalize Config = tuple[Hashable, tuple[Any, ...]] @@ -15,7 +15,7 @@ def recognizes_vpa(vpa: VisiblyPushdownAutomaton, word: Sequence[Any]) -> bool: """Return whether ``vpa`` accepts ``word``. - Simulates the normalized form (:func:`~sofic.automata.vpa_constructions.normalize`): + Simulates the normalized form (:func:`~sofic.automata.vpa.operations.normalize`): the stack holds only pushed symbols, its emptiness stands for the explicit bottom, and a return guarded by ``BOTTOM`` fires only on the empty stack. """ diff --git a/sofic/generators/__init__.py b/sofic/generators/__init__.py index d0866e3..92ce391 100644 --- a/sofic/generators/__init__.py +++ b/sofic/generators/__init__.py @@ -16,7 +16,6 @@ synergistic_information_flow, transfer_entropy, ) -from sofic.generators.epsilon_inference import cssr, spectral, subtree_merge, suggest_lmax from sofic.generators.epsilon_machine import EpsilonMachine from sofic.generators.epsilon_transducer import EpsilonTransducer from sofic.generators.lumping import LumpabilityError, is_lumpable, lump, normalize_partition @@ -38,12 +37,6 @@ from sofic.generators.pfa import ProbabilisticFiniteAutomaton from sofic.generators.quasi_realization import QuasiRealization from sofic.generators.stack_hmm import HiddenMarkovStackModel -from sofic.generators.stack_inference import ( - fit_stack_hmm_mle, - learn_stack_hmm_papni, - stack_cssr, - stack_subtree_merge, -) from sofic.generators.topological_epsilon_enumeration import ( count_topological_epsilon_machines, epsilon_machine_to_idfa_string, @@ -76,17 +69,9 @@ "LumpabilityError", "channel_statistical_complexity", "driven_entropy_rate", - "cssr", "is_lumpable", "lump", "normalize_partition", - "spectral", - "subtree_merge", - "suggest_lmax", - "fit_stack_hmm_mle", - "learn_stack_hmm_papni", - "stack_cssr", - "stack_subtree_merge", "directed_information", "independent_pair_generator", "intrinsic_information_flow", diff --git a/sofic/generators/_suffix_counts.py b/sofic/generators/_suffix_counts.py deleted file mode 100644 index ed3779e..0000000 --- a/sofic/generators/_suffix_counts.py +++ /dev/null @@ -1,21 +0,0 @@ -"""Sliding-window suffix scan shared by the CSSR history counters.""" - -from __future__ import annotations - -from collections.abc import Iterator, Sequence -from typing import Any - - -def infer_alphabet(tokens: Sequence[Any], alphabet: Sequence[Any] | None) -> tuple[Any, ...]: - """``alphabet`` as a tuple, or the distinct ``tokens`` sorted by ``repr``.""" - return tuple(sorted(set(tokens), key=repr)) if alphabet is None else tuple(alphabet) - - -def iter_suffixes(tokens: tuple[Any, ...], max_length: int) -> Iterator[tuple[int, tuple[Any, ...]]]: - """Yield ``(t, tokens[t - L : t])`` for each position ``t`` and each ``L`` in ``0..min(t, max_length)``. - - These are the pasts, up to ``max_length`` long, that precede the token at ``t``. - """ - for t in range(len(tokens)): - for length in range(min(t, max_length) + 1): - yield t, tokens[t - length : t] diff --git a/sofic/generators/base.py b/sofic/generators/base.py index 47325b9..6b37412 100644 --- a/sofic/generators/base.py +++ b/sofic/generators/base.py @@ -103,39 +103,39 @@ def __init__(self, observation_alphabet: frozenset[Any] | None = None, **kwargs: self.observation_alphabet = observation_alphabet if observation_alphabet is not None else frozenset() def sample(self, n: int, rng: np.random.Generator | None = None) -> tuple[list[Any], list[Hashable]]: - from sofic.generators.hmm_inference import sample + from sofic.generators.sampling import sample return sample(self, n, rng) def log_likelihood(self, observations: Sequence[Any]) -> float: - from sofic.generators.hmm_inference import log_likelihood + from sofic.inference.hmm import log_likelihood return log_likelihood(self, observations) def forward(self, observations: Sequence[Any], *, scaled: bool = False) -> np.ndarray: - from sofic.generators.hmm_inference import forward + from sofic.inference.hmm import forward return forward(self, observations, scaled=scaled) def backward(self, observations: Sequence[Any], *, scaled: bool = False) -> np.ndarray: - from sofic.generators.hmm_inference import backward + from sofic.inference.hmm import backward return backward(self, observations, scaled=scaled) def viterbi(self, observations: Sequence[Any]) -> list[Hashable]: - from sofic.generators.hmm_inference import viterbi + from sofic.inference.hmm import viterbi return viterbi(self, observations) def smooth(self, observations: Sequence[Any]) -> np.ndarray: """Return fixed-interval smoothed marginals ``gamma[t, s]``.""" - from sofic.generators.hmm_inference import smooth + from sofic.inference.hmm import smooth return smooth(self, observations) def two_slice_marginals(self, observations: Sequence[Any]) -> np.ndarray: """Return two-slice smoothed marginals ``xi[t, i, j]``.""" - from sofic.generators.hmm_inference import two_slice_marginals + from sofic.inference.hmm import two_slice_marginals return two_slice_marginals(self, observations) @@ -151,9 +151,9 @@ def baum_welch( ) -> tuple[MealyHMM, list[float]]: """Fit parameters by Baum-Welch EM, returning ``(fitted_model, loglik_trace)``. - ``n_restarts`` and ``rng`` are as in :func:`~sofic.generators.hmm_inference.baum_welch`. + ``n_restarts`` and ``rng`` are as in :func:`~sofic.inference.hmm.baum_welch`. """ - from sofic.generators.hmm_inference import baum_welch + from sofic.inference.hmm import baum_welch return baum_welch( self, @@ -167,19 +167,19 @@ def baum_welch( def score(self, observations: Sequence[Any]) -> dict[tuple[Hashable, Any, Hashable], float]: """Return the log-likelihood gradient (Fisher identity) over edge parameters.""" - from sofic.generators.hmm_inference import score + from sofic.inference.hmm import score return score(self, observations) def observed_information(self, observations: Sequence[Any]) -> np.ndarray: """Return the observed information matrix (Louis' identity).""" - from sofic.generators.hmm_inference import observed_information + from sofic.inference.hmm import observed_information return observed_information(self, observations) def standard_errors(self, observations: Sequence[Any]) -> dict[tuple[Hashable, Any, Hashable], float]: """Return asymptotic standard errors of the free edge parameters.""" - from sofic.generators.hmm_inference import standard_errors + from sofic.inference.hmm import standard_errors return standard_errors(self, observations) diff --git a/sofic/generators/epsilon_machine.py b/sofic/generators/epsilon_machine.py index 0372e26..0aa67ee 100644 --- a/sofic/generators/epsilon_machine.py +++ b/sofic/generators/epsilon_machine.py @@ -66,20 +66,20 @@ def from_sequence( ``"spectral"`` for Hankel-SVD learning followed by mixed-state extraction. **kwargs - Forwarded to :func:`~sofic.generators.epsilon_inference.cssr`, - :func:`~sofic.generators.epsilon_inference.subtree_merge`, or - :func:`~sofic.generators.epsilon_inference.spectral`. + Forwarded to :func:`~sofic.inference.cssr.cssr`, + :func:`~sofic.inference.cssr.subtree_merge`, or + :func:`~sofic.inference.spectral.spectral`. """ if method == "cssr": - from sofic.generators.epsilon_inference import cssr + from sofic.inference.cssr.process import cssr return cssr(sequence, **kwargs) if method == "subtree": - from sofic.generators.epsilon_inference import subtree_merge + from sofic.inference.cssr.subtree import subtree_merge return subtree_merge(sequence, **kwargs) if method == "spectral": - from sofic.generators.epsilon_inference import spectral + from sofic.inference.spectral import spectral return spectral(sequence, **kwargs) raise ValueError(f"unknown inference method {method!r}") diff --git a/sofic/generators/epsilon_transducer.py b/sofic/generators/epsilon_transducer.py index c955e1c..e6406ad 100644 --- a/sofic/generators/epsilon_transducer.py +++ b/sofic/generators/epsilon_transducer.py @@ -128,7 +128,7 @@ def from_paired_sequences( **kwargs: Any, ) -> EpsilonTransducer: """Reconstruct an ε-transducer from paired input/output sequences via transCSSR.""" - from sofic.generators.epsilon_transducer_inference import transcssr + from sofic.inference.cssr.transducer import transcssr return transcssr(inputs, outputs, **kwargs) diff --git a/sofic/generators/hmm_inference.py b/sofic/generators/hmm_inference.py deleted file mode 100644 index 8674f4e..0000000 --- a/sofic/generators/hmm_inference.py +++ /dev/null @@ -1,688 +0,0 @@ -"""Inference for hidden Markov models. - -Forward/backward/Viterbi decoding and sampling, plus the Cappe, Moulines & -Ryden (2005) toolbox: fixed-interval smoothing (one- and two-slice marginals), -Baum-Welch EM parameter re-estimation, and the score / observed information via -the Fisher and Louis identities. -""" - -from __future__ import annotations - -import warnings -from collections import defaultdict -from collections.abc import Hashable, Iterable, Sequence -from typing import Any - -import numpy as np - -from sofic.generators.base import HiddenMarkovModel -from sofic.generators.matrices import emission_tensors, symbol_matrices - - -def _as_mealy_hmm(hmm: HiddenMarkovModel) -> Any: - """Return a Mealy-style representation through the HMM representation hook.""" - return hmm.to_mealy() - - -def _forward_scaled( - pi: np.ndarray, - joint: dict[Any, np.ndarray], - obs: list[Any], -) -> tuple[np.ndarray, np.ndarray]: - """Return per-step-normalized forward messages and log scaling factors. - - ``alpha_hat[t]`` sums to one; ``log2 P(obs) = log_scales.sum()``. A ``-inf`` - entry in ``log_scales`` marks an impossible step. Normalizing each step avoids - the underflow that makes the raw forward product vanish for long sequences. - """ - n = len(pi) - alpha_hat = np.zeros((len(obs) + 1, n), dtype=float) - log_scales = np.zeros(len(obs) + 1, dtype=float) - total0 = float(pi.sum()) - if total0 <= 0.0: - log_scales[0] = -np.inf - return alpha_hat, log_scales - alpha_hat[0] = pi / total0 - log_scales[0] = float(np.log2(total0)) - for t, symbol in enumerate(obs): - matrix = joint.get(symbol) - if matrix is None: - log_scales[t + 1] = -np.inf - continue - row = alpha_hat[t] @ matrix - scale = float(row.sum()) - if scale <= 0.0: - log_scales[t + 1] = -np.inf - continue - alpha_hat[t + 1] = row / scale - log_scales[t + 1] = float(np.log2(scale)) - return alpha_hat, log_scales - - -def forward(hmm: HiddenMarkovModel, observations: Sequence[Any], *, scaled: bool = False) -> np.ndarray: - """Return forward messages ``alpha[t, s]`` for ``len(observations)+1`` rows. - - With ``scaled=True`` each row is normalized to sum to one (the numerically - stable message used for posteriors); otherwise the raw messages are returned. - """ - pi, joint = emission_tensors(hmm) - obs = list(observations) - if scaled: - alpha_hat, _log_scales = _forward_scaled(pi, joint, obs) - return alpha_hat - n = len(pi) - alpha = np.zeros((len(obs) + 1, n), dtype=float) - alpha[0] = pi - for t, symbol in enumerate(obs): - matrix = joint.get(symbol) - if matrix is None: - alpha[t + 1] = 0.0 - else: - alpha[t + 1] = alpha[t] @ matrix - return alpha - - -def backward(hmm: HiddenMarkovModel, observations: Sequence[Any], *, scaled: bool = False) -> np.ndarray: - """Return backward messages ``beta[t, s]`` for ``len(observations)+1`` rows. - - With ``scaled=True`` each row is normalized to sum to one. The smoothed - posterior is then ``normalize(alpha_hat[t] * beta_hat[t])`` (the per-row - scaling constants cancel on renormalization). - """ - joint = symbol_matrices(_as_mealy_hmm(hmm)) - n = next(iter(joint.values())).shape[0] if joint else len(_as_mealy_hmm(hmm).reindex()) - obs = list(observations) - beta = np.zeros((len(obs) + 1, n), dtype=float) - beta[len(obs)] = 1.0 - for t in range(len(obs) - 1, -1, -1): - matrix = joint.get(obs[t]) - if matrix is None: - beta[t] = 0.0 - else: - beta[t] = matrix @ beta[t + 1] - if scaled: - total = float(beta[t].sum()) - if total > 0.0: - beta[t] = beta[t] / total - return beta - - -def _backward_scaled(joint: dict[Any, np.ndarray], obs: list[Any], n_states: int) -> np.ndarray: - """Per-row-normalized backward messages from precomputed transition tensors. - - ``beta_hat[t]`` sums to one; the per-row scaling constants cancel against the - forward scaling when the smoothed posterior is renormalized. Shares tensors - with the forward pass so smoothing and EM avoid recomputing them. - """ - beta = np.zeros((len(obs) + 1, n_states), dtype=float) - beta[len(obs)] = 1.0 - for t in range(len(obs) - 1, -1, -1): - matrix = joint.get(obs[t]) - row = beta[t + 1] if matrix is None else matrix @ beta[t + 1] - beta[t] = 0.0 if matrix is None else row - total = float(beta[t].sum()) - if total > 0.0: - beta[t] = beta[t] / total - return beta - - -def log_likelihood(hmm: HiddenMarkovModel, observations: Sequence[Any]) -> float: - """Log-likelihood ``log2 P(observations)`` in bits. - - Uses the per-step-scaled forward recursion so the result stays finite for long - sequences instead of underflowing to ``-inf``. - """ - pi, joint = emission_tensors(hmm) - _alpha_hat, log_scales = _forward_scaled(pi, joint, list(observations)) - if not np.all(np.isfinite(log_scales)): - return float("-inf") - return float(log_scales.sum()) - - -def smooth(hmm: HiddenMarkovModel, observations: Sequence[Any]) -> np.ndarray: - r"""Return fixed-interval smoothed marginals ``gamma[t, s]``. - - ``gamma[t, s] = P(X_t = s \mid Y_{0:n-1})`` for ``t = 0, ..., n`` (there are - ``n + 1`` hidden states behind ``n`` edge emissions). Computed as the - per-row-renormalized product of the scaled forward and backward messages, the - forward-backward smoother of Cappe, Moulines & Ryden (2005, Section 3.2). - Rows for observation sequences of zero probability are returned as zeros. - """ - pi, joint = emission_tensors(hmm) - obs = list(observations) - n_states = len(pi) - alpha_hat, log_scales = _forward_scaled(pi, joint, obs) - if not np.all(np.isfinite(log_scales)): - return np.zeros((len(obs) + 1, n_states), dtype=float) - beta_hat = _backward_scaled(joint, obs, n_states) - gamma = alpha_hat * beta_hat - row_sums = gamma.sum(axis=1, keepdims=True) - with np.errstate(invalid="ignore", divide="ignore"): - gamma = np.where(row_sums > 0.0, gamma / row_sums, 0.0) - return gamma - - -def two_slice_marginals(hmm: HiddenMarkovModel, observations: Sequence[Any]) -> np.ndarray: - r"""Return two-slice smoothed marginals ``xi[t, i, j]``. - - ``xi[t, i, j] = P(X_t = i, X_{t+1} = j \mid Y_{0:n-1})`` for ``t = 0, ..., n-1``, - where the transition at index ``t`` emits ``Y_t`` (Cappe, Moulines & Ryden, - 2005, Section 3.2). Marginalizing over ``j`` recovers ``gamma[t]`` for - ``t < n``. Returns an all-zero tensor for zero-probability sequences. - """ - pi, joint = emission_tensors(hmm) - obs = list(observations) - n_states = len(pi) - xi = np.zeros((len(obs), n_states, n_states), dtype=float) - alpha_hat, log_scales = _forward_scaled(pi, joint, obs) - if not np.all(np.isfinite(log_scales)): - return xi - beta_hat = _backward_scaled(joint, obs, n_states) - for t, symbol in enumerate(obs): - matrix = joint.get(symbol) - if matrix is None: - continue - block = alpha_hat[t][:, None] * matrix * beta_hat[t + 1][None, :] - total = float(block.sum()) - if total > 0.0: - xi[t] = block / total - return xi - - -def _expected_edge_counts( - pi: np.ndarray, - joint: dict[Any, np.ndarray], - obs: list[Any], -) -> tuple[dict[tuple[int, Any, int], float], np.ndarray, np.ndarray, float]: - r"""Expected sufficient statistics for one observation sequence. - - Returns ``(edge_counts, source_totals, gamma0, loglik)`` where - - - ``edge_counts[(i, symbol, j)]`` is - :math:`\sum_t P(X_t = i, Y_t = symbol, X_{t+1} = j \mid Y)`, the expected - number of uses of edge ``i --symbol--> j``; - - ``source_totals[i] = \sum_{t=0}^{n-1} P(X_t = i \mid Y)`` is the expected - number of transitions out of state ``i`` (the Baum-Welch denominator); - - ``gamma0`` is the smoothed marginal of the initial state ``X_0``; - - ``loglik`` is the log-likelihood of the sequence in bits. - - Only edges present in ``joint`` (structural support) receive mass, so the - statistics preserve the model topology. - """ - n_states = len(pi) - alpha_hat, log_scales = _forward_scaled(pi, joint, obs) - if not np.all(np.isfinite(log_scales)): - return {}, np.zeros(n_states), np.zeros(n_states), float("-inf") - beta_hat = _backward_scaled(joint, obs, n_states) - edge_counts: dict[tuple[int, Any, int], float] = {} - source_totals = np.zeros(n_states, dtype=float) - for t, symbol in enumerate(obs): - matrix = joint.get(symbol) - if matrix is None: - continue - block = alpha_hat[t][:, None] * matrix * beta_hat[t + 1][None, :] - total = float(block.sum()) - if total <= 0.0: - continue - block = block / total - source_totals += block.sum(axis=1) - for i, j in np.argwhere(block > 0.0): - key = (int(i), symbol, int(j)) - edge_counts[key] = edge_counts.get(key, 0.0) + float(block[i, j]) - g0 = alpha_hat[0] * beta_hat[0] - s0 = float(g0.sum()) - gamma0 = g0 / s0 if s0 > 0.0 else np.zeros(n_states) - return edge_counts, source_totals, gamma0, float(log_scales.sum()) - - -def _as_sequence_list(sequences: Iterable[Any]) -> list[list[Any]]: - """Normalize ``sequences`` to a list of observation sequences. - - Accepts either a single flat observation sequence (e.g. ``[0, 1, 0]``) or an - iterable of sequences (e.g. ``[[0, 1], [1, 0]]``). A single sequence is - detected when its first element is not itself a non-string sequence. - """ - seqs = list(sequences) - if not seqs: - return [] - first = seqs[0] - if isinstance(first, (list, tuple)) and not isinstance(first, (str, bytes)): - return [list(seq) for seq in seqs] - return [seqs] - - -def baum_welch( - hmm: HiddenMarkovModel, - sequences: Iterable[Any], - *, - max_iter: int = 100, - tol: float = 1e-6, - estimate_initial: bool = True, - n_restarts: int = 1, - rng: np.random.Generator | int | None = None, - return_restarts: bool = False, -) -> tuple[Any, list[float]] | tuple[Any, list[float], list[float]]: - r"""Fit HMM parameters by Baum-Welch (EM) expectation-maximization. - - Re-estimates the Mealy joint edge law - :math:`A_o[i, j] = P(X_{t+1} = j, O = o \mid X_t = i)` and (optionally) the - initial distribution from data, holding the transition-graph topology fixed: - structurally absent edges receive zero expected count and stay absent, so the - fitted model generates the same sofic shift as ``hmm``. This is the EM - algorithm for probabilistic functions of finite Markov chains of Baum, Petrie, - Soules & Weiss and Cappe, Moulines & Ryden (2005, Chapter 10); see also - Rabiner (1989). - - ``sequences`` may be a single observation sequence or an iterable of - sequences (several sequences are needed to identify the initial distribution; - Cappe, Moulines & Ryden, 2005, Section 10.1). Unifilarity is *not* preserved, - so the fit is returned as a plain :class:`~sofic.generators.mealy.MealyHMM`. - - Returns ``(fitted_model, loglik_trace)`` where ``loglik_trace`` is the - non-decreasing sequence of total log-likelihoods (bits) observed before each - parameter update. - - EM converges to a local maximum of the likelihood. With ``n_restarts > 1`` the - first run starts from ``hmm``'s parameters and each further run from edge laws - drawn uniformly (Dirichlet(1)) over each state's structurally allowed edges; - the fit with the highest final log-likelihood is returned. Pass - ``return_restarts=True`` to also get every run's final log-likelihood, which - shows whether near-equal optima exist. - """ - from sofic.generators.mealy import MealyHMM - - mealy = hmm.to_mealy() - idx = mealy.reindex() - n_states = len(idx) - states = [idx.state(i) for i in range(n_states)] - alphabet = frozenset(mealy.observation_alphabet) - seqs = _as_sequence_list(sequences) - - pi, joint = emission_tensors(mealy) - support = { - (i, symbol, j) - for symbol, matrix in joint.items() - for i in range(n_states) - for j in range(n_states) - if matrix[i, j] > 0.0 - } - - if n_restarts < 1: - raise ValueError("n_restarts must be at least 1") - generator = rng if isinstance(rng, np.random.Generator) else np.random.default_rng(rng) - runs = [] - for restart in range(n_restarts): - start_joint = joint if restart == 0 else _random_edge_law(joint, support, n_states, generator) - runs.append( - _baum_welch_run(pi, start_joint, seqs, max_iter=max_iter, tol=tol, estimate_initial=estimate_initial) - ) - finals = [trace[-1] if trace else float("-inf") for _pi, _joint, trace in runs] - pi, joint, loglik_trace = runs[int(np.argmax(finals))] - - fitted = MealyHMM( - initial_distribution={states[i]: float(pi[i]) for i in range(n_states) if pi[i] > 0.0}, - observation_alphabet=alphabet, - ) - for state in states: - fitted.graph.add_state(state) - for i, symbol, j in sorted(support, key=lambda edge: (edge[0], str(edge[1]), edge[2])): - prob = float(joint[symbol][i, j]) - if prob > 0.0: - fitted.add_transition(states[i], states[j], symbol, prob) - fitted.validate() - if return_restarts: - return fitted, loglik_trace, finals - return fitted, loglik_trace - - -def _random_edge_law( - joint: dict[Any, np.ndarray], - support: set[tuple[int, Any, int]], - n_states: int, - rng: np.random.Generator, -) -> dict[Any, np.ndarray]: - """Edge laws drawn uniformly over each state's allowed ``(symbol, target)`` edges.""" - new_joint = {symbol: np.zeros((n_states, n_states), dtype=float) for symbol in joint} - for i in range(n_states): - edges = sorted(((symbol, j) for (source, symbol, j) in support if source == i), key=lambda e: (str(e[0]), e[1])) - if not edges: - continue - weights = rng.dirichlet(np.ones(len(edges))) - for (symbol, j), weight in zip(edges, weights, strict=True): - new_joint[symbol][i, j] = weight - return new_joint - - -def _baum_welch_run( - pi: np.ndarray, - joint: dict[Any, np.ndarray], - seqs: list[Any], - *, - max_iter: int, - tol: float, - estimate_initial: bool, -) -> tuple[np.ndarray, dict[Any, np.ndarray], list[float]]: - """One EM run from ``(pi, joint)``; returns the final parameters and trace.""" - n_states = len(pi) - loglik_trace: list[float] = [] - prev_ll: float | None = None - for _iteration in range(max_iter): - total_edge_counts: dict[tuple[int, Any, int], float] = defaultdict(float) - total_source = np.zeros(n_states, dtype=float) - gamma0_sum = np.zeros(n_states, dtype=float) - total_ll = 0.0 - skipped = 0 - for obs in seqs: - edge_counts, source_totals, gamma0, loglik = _expected_edge_counts(pi, joint, obs) - if not np.isfinite(loglik): - skipped += 1 - continue - for key, value in edge_counts.items(): - total_edge_counts[key] += value - total_source += source_totals - gamma0_sum += gamma0 - total_ll += loglik - if seqs and skipped == len(seqs): - raise ValueError("every observation sequence has zero probability under the model") - if skipped and _iteration == 0: - warnings.warn( - f"{skipped} of {len(seqs)} sequences have zero probability under the model and are ignored", - RuntimeWarning, - stacklevel=3, - ) - loglik_trace.append(total_ll) - if prev_ll is not None and abs(total_ll - prev_ll) < tol: - break - prev_ll = total_ll - - new_joint = {symbol: np.zeros((n_states, n_states), dtype=float) for symbol in joint} - for (i, symbol, j), count in total_edge_counts.items(): - if total_source[i] > 0.0: - new_joint[symbol][i, j] = count / total_source[i] - for i in range(n_states): - if total_source[i] <= 0.0: - for symbol in joint: - new_joint[symbol][i, :] = joint[symbol][i, :] - joint = new_joint - if estimate_initial: - mass = float(gamma0_sum.sum()) - if mass > 0.0: - pi = gamma0_sum / mass - return pi, joint, loglik_trace - - -def score(hmm: HiddenMarkovModel, observations: Sequence[Any]) -> dict[tuple[Hashable, Any, Hashable], float]: - r"""Return the score (gradient of the log-likelihood) via the Fisher identity. - - For each edge ``i --o--> j``, returns - :math:`\partial \log P(Y) / \partial A_o[i, j] = E[N_{i,o,j} \mid Y] / A_o[i, j]`, - where ``N`` is the (unobserved) edge-use count. This is Fisher's identity, - ``\nabla \log L(\theta) = E[\nabla \log f(X, Y; \theta) \mid Y]`` (Cappe, - Moulines & Ryden, 2005, Section 10.2.3), evaluated in the raw (unconstrained) - joint-edge parameters. Keys are ``(source, symbol, target)`` state labels. - - The score is the gradient of the *natural* log-likelihood - :math:`\ln P(Y) = \ln 2 \cdot` :func:`log_likelihood`, the convention under - which the observed information and standard errors are standard. - """ - mealy = hmm.to_mealy() - idx = mealy.reindex() - pi, joint = emission_tensors(mealy) - edge_counts, _source_totals, _gamma0, loglik = _expected_edge_counts(pi, joint, list(observations)) - if not np.isfinite(loglik): - raise ValueError("observations have zero probability under the model; score is undefined") - result: dict[tuple[Hashable, Any, Hashable], float] = {} - n_states = len(pi) - for symbol, matrix in joint.items(): - for i in range(n_states): - for j in range(n_states): - prob = float(matrix[i, j]) - if prob > 0.0: - count = edge_counts.get((i, symbol, j), 0.0) - result[(idx.state(i), symbol, idx.state(j))] = count / prob - return result - - -def _free_parameterization( - joint: dict[Any, np.ndarray], - n_states: int, -) -> tuple[list[tuple[int, Any, int]], list[tuple[int, Any, int]], list[int]]: - """Build the free multinomial parameterization of the joint edge law. - - Each source state whose outgoing edges number ``k >= 2`` contributes ``k - 1`` - free parameters (its last edge in canonical order is the reference). Returns - ``(free_edges, reference_by_param, source_by_param)``: the edge for each free - parameter, the reference edge of its source block, and the source-state index. - """ - free_edges: list[tuple[int, Any, int]] = [] - reference_by_param: list[tuple[int, Any, int]] = [] - source_by_param: list[int] = [] - for i in range(n_states): - out_edges = sorted( - ((i, symbol, j) for symbol, matrix in joint.items() for j in range(n_states) if matrix[i, j] > 0.0), - key=lambda edge: (str(edge[1]), edge[2]), - ) - if len(out_edges) < 2: - continue - reference = out_edges[-1] - for edge in out_edges[:-1]: - free_edges.append(edge) - reference_by_param.append(reference) - source_by_param.append(i) - return free_edges, reference_by_param, source_by_param - - -def observed_information(hmm: HiddenMarkovModel, observations: Sequence[Any]) -> np.ndarray: - r"""Return the observed information matrix via Louis' identity. - - The observed information ``J = -\partial^2 \log L / \partial\theta^2`` for the - free multinomial parameters of the joint edge law is obtained from Louis' - (1982) identity, - - .. math:: J = E[-\partial^2 \ell_c \mid Y] - \operatorname{Cov}(\partial \ell_c \mid Y), - - where :math:`\ell_c` is the complete-data log-likelihood (Cappe, Moulines & - Ryden, 2005, Section 10.2.3). The complete-data information ``B`` follows from - the expected edge counts; the conditional covariance of the complete-data - score is computed exactly by a forward smoothing recursion for the first and - second moments of the additive score functional. The matrix is ordered by - :func:`free_parameter_labels`; an empty ``(0, 0)`` matrix is returned when the - model has no free parameters. - """ - mealy = hmm.to_mealy() - pi, joint = emission_tensors(mealy) - obs = list(observations) - n_states = len(pi) - - free_edges, reference_by_param, source_by_param = _free_parameterization(joint, n_states) - d = len(free_edges) - if d == 0: - return np.zeros((0, 0), dtype=float) - - edge_counts, _source_totals, _gamma0, loglik = _expected_edge_counts(pi, joint, obs) - if not np.isfinite(loglik): - raise ValueError("observations have zero probability under the model; information is undefined") - - prob_of = {edge: float(joint[edge[1]][edge[0], edge[2]]) for edge in set(free_edges) | set(reference_by_param)} - - # Complete-data information B = E[-d^2 l_c | Y], block-diagonal by source state. - complete_information = np.zeros((d, d), dtype=float) - for p in range(d): - ref_p = reference_by_param[p] - count_ref = edge_counts.get(ref_p, 0.0) - ref_term = count_ref / prob_of[ref_p] ** 2 - for q in range(d): - if source_by_param[p] != source_by_param[q]: - continue - value = ref_term - if p == q: - edge_p = free_edges[p] - value += edge_counts.get(edge_p, 0.0) / prob_of[edge_p] ** 2 - complete_information[p, q] = value - - # Per-transition score contribution s(edge) as a d-vector (sparse per source block). - edge_score: dict[tuple[int, Any, int], np.ndarray] = {} - for p, edge in enumerate(free_edges): - edge_score.setdefault(edge, np.zeros(d))[p] += 1.0 / prob_of[edge] - for p, ref in enumerate(reference_by_param): - edge_score.setdefault(ref, np.zeros(d))[p] += -1.0 / prob_of[ref] - zero_d = np.zeros(d) - - # Forward smoothing recursion for E[S | Y] and E[S S^T | Y] of the additive - # complete-data score functional S = sum_t s(edge_t). - alpha_hat, _log_scales = _forward_scaled(pi, joint, obs) - first = np.zeros((n_states, d), dtype=float) - second = np.zeros((n_states, d, d), dtype=float) - for t, symbol in enumerate(obs): - matrix = joint.get(symbol) - if matrix is None: - continue - weight = alpha_hat[t][:, None] * matrix # weight[i, k] = P(X_t=i, X_{t+1}=k, Y_t | Y_{0:t-1}) - denom = weight.sum(axis=0) - new_first = np.zeros((n_states, d), dtype=float) - new_second = np.zeros((n_states, d, d), dtype=float) - for k in range(n_states): - if denom[k] <= 0.0: - continue - for i in range(n_states): - if weight[i, k] <= 0.0: - continue - retro = weight[i, k] / denom[k] # P(X_t=i | X_{t+1}=k, Y_{0:t}) - s_vec = edge_score.get((i, symbol, k), zero_d) - first_i = first[i] - combined = first_i + s_vec - new_first[k] += retro * combined - cross = np.outer(first_i, s_vec) - new_second[k] += retro * (second[i] + cross + cross.T + np.outer(s_vec, s_vec)) - first, second = new_first, new_second - - phi_final = alpha_hat[len(obs)] - expected_score = phi_final @ first - expected_outer = np.einsum("k,kpq->pq", phi_final, second) - score_covariance = expected_outer - np.outer(expected_score, expected_score) - return complete_information - score_covariance - - -def free_parameter_labels(hmm: HiddenMarkovModel) -> list[tuple[Hashable, Any, Hashable]]: - """Return the ``(source, symbol, target)`` label for each free parameter. - - The order matches the rows and columns of :func:`observed_information` and the - entries of :func:`standard_errors`. - """ - mealy = hmm.to_mealy() - idx = mealy.reindex() - joint = symbol_matrices(mealy) - free_edges, _reference, _source = _free_parameterization(joint, len(idx)) - return [(idx.state(i), symbol, idx.state(j)) for i, symbol, j in free_edges] - - -def standard_errors( - hmm: HiddenMarkovModel, - observations: Sequence[Any], -) -> dict[tuple[Hashable, Any, Hashable], float]: - r"""Return asymptotic standard errors of the free edge parameters. - - Standard errors are ``sqrt(diag(J^{-1}))`` where ``J`` is the - :func:`observed_information` matrix (Cappe, Moulines & Ryden, 2005, - Section 10.2.3). Uses the Moore-Penrose pseudoinverse when ``J`` is singular; - a non-positive variance estimate (numerically unidentified parameter) yields - ``nan``. Keyed by the labels from :func:`free_parameter_labels`. - """ - labels = free_parameter_labels(hmm) - information = observed_information(hmm, observations) - if information.shape[0] == 0: - return {} - try: - covariance = np.linalg.inv(information) - except np.linalg.LinAlgError: - covariance = np.linalg.pinv(information) - variances = np.diag(covariance) - with np.errstate(invalid="ignore"): - errors = np.where(variances > 0.0, np.sqrt(variances), np.nan) - return dict(zip(labels, (float(value) for value in errors), strict=True)) - - -def _log_probabilities(values: np.ndarray) -> np.ndarray: - log_values = np.full(values.shape, -np.inf, dtype=float) - positive = values > 0.0 - log_values[positive] = np.log(values[positive]) - return log_values - - -def viterbi(hmm: HiddenMarkovModel, observations: Sequence[Any]) -> list[Hashable]: - mealy = _as_mealy_hmm(hmm) - idx = mealy.reindex() - pi, joint = emission_tensors(mealy) - n = len(idx) - obs = list(observations) - if n == 0: - return [] - if not obs: - if not np.any(pi > 0.0): - return [] - return [idx.state(int(np.argmax(pi)))] - - log_pi = _log_probabilities(pi) - viterbi_log = np.full((len(obs), n), -np.inf, dtype=float) - backpointer = np.full((len(obs), n), -1, dtype=int) - - matrix0 = joint.get(obs[0]) - if matrix0 is not None: - log_matrix0 = _log_probabilities(matrix0) - for j in range(n): - best = log_pi + log_matrix0[:, j] - viterbi_log[0, j] = np.max(best) - backpointer[0, j] = int(np.argmax(best)) - - for t in range(1, len(obs)): - matrix = joint.get(obs[t]) - if matrix is None: - continue - log_matrix = _log_probabilities(matrix) - for j in range(n): - scores = viterbi_log[t - 1] + log_matrix[:, j] - viterbi_log[t, j] = np.max(scores) - backpointer[t, j] = int(np.argmax(scores)) - - if not np.any(np.isfinite(viterbi_log[-1])): - return [] - - path = [0] * len(obs) - path[-1] = int(np.argmax(viterbi_log[-1])) - for t in range(len(obs) - 2, -1, -1): - path[t] = backpointer[t + 1, path[t + 1]] - return [idx.state(i) for i in path] - - -def sample( - hmm: HiddenMarkovModel, - n: int, - rng: np.random.Generator | None = None, -) -> tuple[list[Any], list[Hashable]]: - generator = rng if rng is not None else np.random.default_rng() - mealy = _as_mealy_hmm(hmm) - idx = mealy.reindex() - pi, joint = emission_tensors(mealy) - total = float(np.sum(pi)) - if not total > 0.0: - raise ValueError("cannot sample: the initial state distribution has no mass") - state = int(generator.choice(len(idx), p=pi / total)) - - observations: list[Any] = [] - states: list[Hashable] = [] - for _ in range(n): - states.append(idx.state(state)) - row_sum = sum(matrix[state].sum() for matrix in joint.values()) - if row_sum <= 0.0: - break - symbol_probs = np.array([joint[sym][state].sum() for sym in joint], dtype=float) - symbol_probs /= symbol_probs.sum() - symbol_index = int(generator.choice(len(joint), p=symbol_probs)) - symbol = list(joint.keys())[symbol_index] - observations.append(symbol) - matrix = joint[symbol] - row = matrix[state] - if row.sum() <= 0.0: - break - state = int(generator.choice(len(idx), p=row / row.sum())) - return observations, states diff --git a/sofic/generators/sampling.py b/sofic/generators/sampling.py new file mode 100644 index 0000000..d11a90a --- /dev/null +++ b/sofic/generators/sampling.py @@ -0,0 +1,45 @@ +"""Sampling observation and hidden-state sequences from hidden Markov models.""" + +from __future__ import annotations + +from collections.abc import Hashable +from typing import Any + +import numpy as np + +from sofic.generators.base import HiddenMarkovModel +from sofic.generators.matrices import emission_tensors + + +def sample( + hmm: HiddenMarkovModel, + n: int, + rng: np.random.Generator | None = None, +) -> tuple[list[Any], list[Hashable]]: + generator = rng if rng is not None else np.random.default_rng() + mealy = hmm.to_mealy() + idx = mealy.reindex() + pi, joint = emission_tensors(mealy) + total = float(np.sum(pi)) + if not total > 0.0: + raise ValueError("cannot sample: the initial state distribution has no mass") + state = int(generator.choice(len(idx), p=pi / total)) + + observations: list[Any] = [] + states: list[Hashable] = [] + for _ in range(n): + states.append(idx.state(state)) + row_sum = sum(matrix[state].sum() for matrix in joint.values()) + if row_sum <= 0.0: + break + symbol_probs = np.array([joint[sym][state].sum() for sym in joint], dtype=float) + symbol_probs /= symbol_probs.sum() + symbol_index = int(generator.choice(len(joint), p=symbol_probs)) + symbol = list(joint.keys())[symbol_index] + observations.append(symbol) + matrix = joint[symbol] + row = matrix[state] + if row.sum() <= 0.0: + break + state = int(generator.choice(len(idx), p=row / row.sum())) + return observations, states diff --git a/sofic/generators/topological_epsilon_enumeration.py b/sofic/generators/topological_epsilon_enumeration.py index 4e24939..72e7e76 100644 --- a/sofic/generators/topological_epsilon_enumeration.py +++ b/sofic/generators/topological_epsilon_enumeration.py @@ -11,7 +11,7 @@ import numpy as np -from sofic.automata.idfa import ( +from sofic.automata.enumeration.idfa import ( MISSING_TRANSITION, IDFAEnumerationError, _delta_table, @@ -128,7 +128,7 @@ def epsilon_machine_to_idfa_string( """Encode an ε-machine as an incomplete accessible DFA transition string. Probabilities are ignored. Missing symbol transitions are encoded with - :data:`sofic.automata.idfa.MISSING_TRANSITION`. If ``canonical`` is true, + :data:`sofic.automata.enumeration.idfa.MISSING_TRANSITION`. If ``canonical`` is true, all states are tried as roots and the rank-minimal IDFA string is returned. Otherwise, the first state in deterministic label order is used as the root. """ diff --git a/sofic/inference/__init__.py b/sofic/inference/__init__.py index 76fdf41..219ab3d 100644 --- a/sofic/inference/__init__.py +++ b/sofic/inference/__init__.py @@ -1,6 +1,21 @@ -"""Inference algorithms for stochastic generators.""" +"""Inference algorithms for stochastic generators. -from sofic.inference import bayesian +The names ``cssr`` and ``spectral`` bound here are the learner functions; they +shadow the :mod:`sofic.inference.cssr` and :mod:`sofic.inference.spectral` +modules as attributes, so import from those modules with ``from ... import``. +""" + +from sofic.inference import bayesian, hmm +from sofic.inference.cssr import ( + cssr, + fit_stack_hmm_mle, + learn_stack_hmm_papni, + stack_cssr, + stack_subtree_merge, + subtree_merge, + suggest_lmax, + transcssr, +) from sofic.inference.diagnostics import ( GoodnessOfFit, StructureStability, @@ -9,6 +24,19 @@ structure_stability, topology_key, ) +from sofic.inference.hmm import ( + backward, + baum_welch, + forward, + free_parameter_labels, + log_likelihood, + observed_information, + score, + smooth, + standard_errors, + two_slice_marginals, + viterbi, +) from sofic.inference.model_selection import ( ModelScores, WAICResult, @@ -28,11 +56,14 @@ project_to_epsilon_machine, project_to_mealy, project_to_nmachine, + spectral, spectral_singular_values, ) __all__ = [ "bayesian", + "hmm", + "cssr", "GoodnessOfFit", "StructureStability", "goodness_of_fit", @@ -55,5 +86,24 @@ "project_to_epsilon_machine", "project_to_mealy", "project_to_nmachine", + "spectral", "spectral_singular_values", + "subtree_merge", + "suggest_lmax", + "transcssr", + "stack_cssr", + "stack_subtree_merge", + "fit_stack_hmm_mle", + "learn_stack_hmm_papni", + "baum_welch", + "viterbi", + "forward", + "backward", + "smooth", + "two_slice_marginals", + "log_likelihood", + "score", + "observed_information", + "standard_errors", + "free_parameter_labels", ] diff --git a/sofic/inference/cssr/__init__.py b/sofic/inference/cssr/__init__.py new file mode 100644 index 0000000..e954d0f --- /dev/null +++ b/sofic/inference/cssr/__init__.py @@ -0,0 +1,54 @@ +"""Causal-State Splitting Reconstruction and its relatives. + +Process CSSR follows Shalizi, Shalizi & Crutchfield (arXiv:cs/0210025); subtree +merging follows Crutchfield & Young (PRL 1989; PRE 1994); transCSSR generalizes +CSSR to ε-transducers (Barnett & Crutchfield, J. Stat. Phys. 161:2 (2015)); stack +CSSR lifts CSSR to (suffix, stack) configurations of hidden Markov stack models. +""" + +from sofic.inference.cssr.counts import ( + ConfigurationHistory, + History, + JointHistory, + JointSuffixCounts, + StackSuffixCounts, + SuffixCounts, +) +from sofic.inference.cssr.process import cssr, suggest_lmax +from sofic.inference.cssr.significance import ( + MorphTest, + TableTest, + aggregates_differ, + morph_test_score, + morphs_differ, +) +from sofic.inference.cssr.stack import ( + fit_stack_hmm_mle, + learn_stack_hmm_papni, + stack_cssr, + stack_subtree_merge, +) +from sofic.inference.cssr.subtree import subtree_merge +from sofic.inference.cssr.transducer import transcssr + +__all__ = [ + "ConfigurationHistory", + "History", + "JointHistory", + "JointSuffixCounts", + "MorphTest", + "StackSuffixCounts", + "SuffixCounts", + "TableTest", + "aggregates_differ", + "cssr", + "fit_stack_hmm_mle", + "learn_stack_hmm_papni", + "morph_test_score", + "morphs_differ", + "stack_cssr", + "stack_subtree_merge", + "subtree_merge", + "suggest_lmax", + "transcssr", +] diff --git a/sofic/inference/cssr/counts.py b/sofic/inference/cssr/counts.py new file mode 100644 index 0000000..bc3c523 --- /dev/null +++ b/sofic/inference/cssr/counts.py @@ -0,0 +1,249 @@ +"""Empirical history counts for the CSSR family. + +:class:`SuffixCounts` counts suffixes of a single process, :class:`JointSuffixCounts` +joint ``(input, output)`` pasts of a channel, and :class:`StackSuffixCounts` +``(suffix, stack)`` configurations of a visibly pushdown process. +""" + +from __future__ import annotations + +from collections import Counter, defaultdict +from collections.abc import Iterable, Iterator, Sequence +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, ClassVar + +if TYPE_CHECKING: + from sofic.automata.learning.papni import DyckAlphabet + +History = tuple[Any, ...] + +JointHistory = tuple[tuple[Any, Any], ...] + +ConfigurationHistory = tuple[tuple[Any, ...], tuple[Any, ...]] + + +def infer_alphabet(tokens: Sequence[Any], alphabet: Sequence[Any] | None) -> tuple[Any, ...]: + """``alphabet`` as a tuple, or the distinct ``tokens`` sorted by ``repr``.""" + return tuple(sorted(set(tokens), key=repr)) if alphabet is None else tuple(alphabet) + + +def iter_suffixes(tokens: tuple[Any, ...], max_length: int) -> Iterator[tuple[int, tuple[Any, ...]]]: + """Yield ``(t, tokens[t - L : t])`` for each position ``t`` and each ``L`` in ``0..min(t, max_length)``. + + These are the pasts, up to ``max_length`` long, that precede the token at ``t``. + """ + for t in range(len(tokens)): + for length in range(min(t, max_length) + 1): + yield t, tokens[t - length : t] + + +@dataclass +class SuffixCounts: + """Empirical counts of histories and following symbols in a sequence.""" + + alphabet: tuple[Any, ...] + history_counts: Counter[History] = field(default_factory=Counter) + next_counts: dict[History, Counter[Any]] = field(default_factory=lambda: defaultdict(Counter)) + + #: History key used as the fallback for an empty history set (overridden by stack counts). + empty_history: ClassVar[History] = () + + @classmethod + def from_sequence( + cls, + sequence: Sequence[Any], + *, + alphabet: Sequence[Any] | None = None, + max_length: int | None = None, + ) -> SuffixCounts: + seq = tuple(sequence) + if not seq: + raise ValueError("sequence must be non-empty") + alphabet = infer_alphabet(seq, alphabet) + unknown = set(seq) - set(alphabet) + if unknown: + raise ValueError(f"symbols {unknown!r} not in alphabet") + counts = cls(alphabet=alphabet) + for t, history in iter_suffixes(seq, max_length if max_length is not None else len(seq)): + counts.history_counts[history] += 1 + counts.next_counts[history][seq[t]] += 1 + return counts + + def morph(self, history: History, *, smoothing: float = 0.0) -> dict[Any, float]: + """MLE (optional additive smoothing) of P(next symbol | history).""" + counts = self.next_counts.get(history, Counter()) + total = sum(counts.values()) + if total == 0: + uniform = 1.0 / len(self.alphabet) + return dict.fromkeys(self.alphabet, uniform) + denom = total + smoothing * len(self.alphabet) + return {symbol: (counts.get(symbol, 0) + smoothing) / denom for symbol in self.alphabet} + + def state_morph(self, histories: set[History], *, smoothing: float = 0.0) -> dict[Any, float]: + """Weighted average of history morphs with weights from occurrence counts.""" + weights = {history: float(self.history_counts.get(history, 0)) for history in histories} + total_weight = sum(weights.values()) + if total_weight <= 0.0: + return self.morph(self.empty_history, smoothing=smoothing) + result = dict.fromkeys(self.alphabet, 0.0) + for history, weight in weights.items(): + morph = self.morph(history, smoothing=smoothing) + for symbol in self.alphabet: + result[symbol] += weight * morph[symbol] + return {symbol: prob / total_weight for symbol, prob in result.items()} + + def marginal_morph(self) -> dict[Any, float]: + """Global next-symbol distribution (IID morph at L=0).""" + counts = Counter() + for _history, counter in self.next_counts.items(): + counts.update(counter) + grand = sum(counts.values()) + if grand == 0: + uniform = 1.0 / len(self.alphabet) + return dict.fromkeys(self.alphabet, uniform) + return {symbol: counts.get(symbol, 0) / grand for symbol in self.alphabet} + + def restricted_to(self, histories: set[History]) -> SuffixCounts: + """Return a plain :class:`SuffixCounts` proxy limited to ``histories``. + + The morph/comparison helpers only read the history sets handed to them, so + stack inference can reuse them by projecting its configuration counts onto a + flat proxy without changing any results. + """ + proxy = SuffixCounts(alphabet=self.alphabet) + proxy.history_counts = Counter({h: self.history_counts.get(h, 0) for h in histories}) + proxy.next_counts = defaultdict(Counter) + for history in histories: + proxy.next_counts[history] = self.next_counts.get(history, Counter()) + return proxy + + +@dataclass +class JointSuffixCounts: + """Empirical counts of joint pasts and following input-conditioned outputs.""" + + input_alphabet: tuple[Any, ...] + output_alphabet: tuple[Any, ...] + history_counts: Counter[JointHistory] = field(default_factory=Counter) + #: ``next_counts[history][input]`` is a Counter over following output symbols. + next_counts: dict[JointHistory, dict[Any, Counter[Any]]] = field(default_factory=dict) + + @classmethod + def from_sequences( + cls, + inputs: Sequence[Any], + outputs: Sequence[Any], + *, + input_alphabet: Sequence[Any] | None = None, + output_alphabet: Sequence[Any] | None = None, + max_length: int, + ) -> JointSuffixCounts: + xs = tuple(inputs) + ys = tuple(outputs) + if len(xs) != len(ys): + raise ValueError("inputs and outputs must have equal length") + if not xs: + raise ValueError("sequences must be non-empty") + counts = cls( + input_alphabet=infer_alphabet(xs, input_alphabet), output_alphabet=infer_alphabet(ys, output_alphabet) + ) + for t, history in iter_suffixes(tuple(zip(xs, ys, strict=True)), max_length): + counts.history_counts[history] += 1 + counts.next_counts.setdefault(history, {}).setdefault(xs[t], Counter())[ys[t]] += 1 + return counts + + def output_counts(self, histories: set[JointHistory], input_symbol: Any) -> Counter[Any]: + observed: Counter[Any] = Counter() + for history in histories: + by_input = self.next_counts.get(history) + if by_input is None: + continue + counter = by_input.get(input_symbol) + if counter is not None: + observed.update(counter) + return observed + + def state_morph(self, histories: set[JointHistory], input_symbol: Any) -> dict[Any, float]: + """Return ``P(output | histories, input_symbol)``.""" + observed = self.output_counts(histories, input_symbol) + total = sum(observed.values()) + if total == 0: + return {} + return {symbol: observed.get(symbol, 0) / total for symbol in self.output_alphabet} + + +#: Aggregated output counts of a state: ``agg[input_symbol]`` is a Counter over outputs. +StateAggregate = dict[Any, Counter[Any]] + + +def _history_aggregate(counts: JointSuffixCounts, history: JointHistory) -> StateAggregate: + return {input_symbol: Counter(counter) for input_symbol, counter in counts.next_counts.get(history, {}).items()} + + +def _merge_aggregate(target: StateAggregate, source: StateAggregate) -> None: + for input_symbol, counter in source.items(): + target.setdefault(input_symbol, Counter()).update(counter) + + +def _state_aggregate(counts: JointSuffixCounts, histories: Iterable[JointHistory]) -> StateAggregate: + aggregate: StateAggregate = {} + for history in histories: + _merge_aggregate(aggregate, _history_aggregate(counts, history)) + return aggregate + + +class StackSuffixCounts(SuffixCounts): + """Empirical counts of (suffix, stack) histories and following symbols. + + Shares the morph / comparison machinery of :class:`SuffixCounts`; only the + empty-history key and the sequence-scanning constructor differ. + """ + + empty_history: ClassVar[History] = ((), ()) + + def __init__( + self, + alphabet: tuple[Any, ...], + history_counts: Counter[ConfigurationHistory] | None = None, + next_counts: dict[ConfigurationHistory, Counter[Any]] | None = None, + ) -> None: + super().__init__( + alphabet=alphabet, + history_counts=history_counts if history_counts is not None else Counter(), + next_counts=next_counts if next_counts is not None else defaultdict(Counter), + ) + + @classmethod + def from_sequence( # type: ignore[override] + cls, + sequence: Sequence[Any], + *, + alphabet: DyckAlphabet, + max_length: int | None = None, + max_stack_depth: int = 8, + ) -> StackSuffixCounts: + seq = tuple(sequence) + if not seq: + raise ValueError("sequence must be non-empty") + stacks: list[tuple[Any, ...]] = [] + stack: tuple[Any, ...] = () + for symbol in seq: + if symbol not in alphabet.symbol_alphabet: + raise ValueError(f"symbol {symbol!r} not in Dyck alphabet") + stacks.append(stack) + stack = _push(stack, symbol, alphabet=alphabet, max_stack_depth=max_stack_depth) + counts = cls(alphabet=infer_alphabet(alphabet.symbol_alphabet, None)) + for t, suffix in iter_suffixes(seq, max_length if max_length is not None else len(seq)): + history = (suffix, stacks[t]) + counts.history_counts[history] += 1 + counts.next_counts[history][seq[t]] += 1 + return counts + + +def _push(stack: tuple[Any, ...], symbol: Any, *, alphabet: DyckAlphabet, max_stack_depth: int) -> tuple[Any, ...]: + """The stack after ``symbol``: calls push (dropping the bottom past ``max_stack_depth``), returns pop.""" + if symbol in alphabet.call_alphabet: + return (*(stack[1:] if len(stack) >= max_stack_depth else stack), symbol) + if symbol in alphabet.return_alphabet: + return stack[:-1] + return stack diff --git a/sofic/generators/epsilon_inference.py b/sofic/inference/cssr/process.py similarity index 55% rename from sofic/generators/epsilon_inference.py rename to sofic/inference/cssr/process.py index 294d868..e75783b 100644 --- a/sofic/generators/epsilon_inference.py +++ b/sofic/inference/cssr/process.py @@ -1,200 +1,27 @@ -"""Sample-based ε-machine reconstruction (CSSR, subtree merging, and spectral). +"""Causal-State Splitting Reconstruction (CSSR) of ε-machines from a sample. -CSSR follows Shalizi, Shalizi & Crutchfield (arXiv:cs/0210025). Subtree merging -follows Crutchfield & Young (PRL 1989; PRE 1994). Spectral reconstruction learns -a weighted finite automaton by Hankel SVD :cite:`Balle2014,Hsu2012` and extracts -causal states as mixed states of the learned operators :cite:`Ellison2009`. +CSSR follows Shalizi, Shalizi & Crutchfield (arXiv:cs/0210025). """ from __future__ import annotations from collections import Counter, defaultdict -from collections.abc import Callable, Iterable, Mapping, Sequence -from dataclasses import dataclass, field -from typing import Any, ClassVar, Literal +from collections.abc import Iterable, Mapping, Sequence +from typing import Any, Literal import numpy as np from sofic.exceptions import StochasticValidationError -from sofic.generators._morph_tests import contingency_table, table_score, table_significant -from sofic.generators._suffix_counts import infer_alphabet, iter_suffixes from sofic.generators.epsilon_machine import EpsilonMachine from sofic.graph import ATTR_EMISSION, ATTR_PROB, TransitionGraph - -History = tuple[Any, ...] - -#: Morph-equality tests: G-test, chi-squared, total-variation threshold, or Monte Carlo exact G-test. -MorphTest = Literal["g", "chi2", "tv", "exact"] - - -@dataclass -class SuffixCounts: - """Empirical counts of histories and following symbols in a sequence.""" - - alphabet: tuple[Any, ...] - history_counts: Counter[History] = field(default_factory=Counter) - next_counts: dict[History, Counter[Any]] = field(default_factory=lambda: defaultdict(Counter)) - - #: History key used as the fallback for an empty history set (overridden by stack counts). - empty_history: ClassVar[History] = () - - @classmethod - def from_sequence( - cls, - sequence: Sequence[Any], - *, - alphabet: Sequence[Any] | None = None, - max_length: int | None = None, - ) -> SuffixCounts: - seq = tuple(sequence) - if not seq: - raise ValueError("sequence must be non-empty") - alphabet = infer_alphabet(seq, alphabet) - unknown = set(seq) - set(alphabet) - if unknown: - raise ValueError(f"symbols {unknown!r} not in alphabet") - counts = cls(alphabet=alphabet) - for t, history in iter_suffixes(seq, max_length if max_length is not None else len(seq)): - counts.history_counts[history] += 1 - counts.next_counts[history][seq[t]] += 1 - return counts - - def morph(self, history: History, *, smoothing: float = 0.0) -> dict[Any, float]: - """MLE (optional additive smoothing) of P(next symbol | history).""" - counts = self.next_counts.get(history, Counter()) - total = sum(counts.values()) - if total == 0: - uniform = 1.0 / len(self.alphabet) - return dict.fromkeys(self.alphabet, uniform) - denom = total + smoothing * len(self.alphabet) - return {symbol: (counts.get(symbol, 0) + smoothing) / denom for symbol in self.alphabet} - - def state_morph(self, histories: set[History], *, smoothing: float = 0.0) -> dict[Any, float]: - """Weighted average of history morphs with weights from occurrence counts.""" - weights = {history: float(self.history_counts.get(history, 0)) for history in histories} - total_weight = sum(weights.values()) - if total_weight <= 0.0: - return self.morph(self.empty_history, smoothing=smoothing) - result = dict.fromkeys(self.alphabet, 0.0) - for history, weight in weights.items(): - morph = self.morph(history, smoothing=smoothing) - for symbol in self.alphabet: - result[symbol] += weight * morph[symbol] - return {symbol: prob / total_weight for symbol, prob in result.items()} - - def marginal_morph(self) -> dict[Any, float]: - """Global next-symbol distribution (IID morph at L=0).""" - counts = Counter() - for _history, counter in self.next_counts.items(): - counts.update(counter) - grand = sum(counts.values()) - if grand == 0: - uniform = 1.0 / len(self.alphabet) - return dict.fromkeys(self.alphabet, uniform) - return {symbol: counts.get(symbol, 0) / grand for symbol in self.alphabet} - - def restricted_to(self, histories: set[History]) -> SuffixCounts: - """Return a plain :class:`SuffixCounts` proxy limited to ``histories``. - - The morph/comparison helpers only read the history sets handed to them, so - stack inference can reuse them by projecting its configuration counts onto a - flat proxy without changing any results. - """ - proxy = SuffixCounts(alphabet=self.alphabet) - proxy.history_counts = Counter({h: self.history_counts.get(h, 0) for h in histories}) - proxy.next_counts = defaultdict(Counter) - for history in histories: - proxy.next_counts[history] = self.next_counts.get(history, Counter()) - return proxy - - -def _observed_counts_for_morph( - counts: SuffixCounts, - histories: set[History], -) -> Counter[Any]: - observed = Counter() - for history in histories: - observed.update(counts.next_counts.get(history, Counter())) - return observed - - -def _contingency_rows( - counts: SuffixCounts, - left_histories: set[History], - right_histories: set[History], -) -> np.ndarray | None: - return contingency_table( - _observed_counts_for_morph(counts, left_histories), - _observed_counts_for_morph(counts, right_histories), - counts.alphabet, - ) - - -def _bonferroni_alpha( - counts: SuffixCounts, - alpha: float, - *, - max_length: int, - min_count: int, - suffix_length: Callable[[History], int] = len, -) -> float: - """``alpha`` divided by the number of suffixes eligible for a split test.""" - eligible = sum( - 1 - for history, following in counts.next_counts.items() - if 0 < suffix_length(history) <= max_length and sum(following.values()) >= max(1, min_count) - ) - return alpha / max(1, eligible) - - -def morphs_differ( - counts: SuffixCounts, - left_histories: set[History], - right_histories: set[History], - *, - alpha: float = 0.05, - test: MorphTest = "g", - delta: float = 0.0, -) -> bool: - """Return whether two history sets have significantly different morphs. - - ``"g"`` is the G-test (log-likelihood ratio) with Yates' continuity correction - when the table has one degree of freedom, i.e. two observed symbols; this is - the statistic ``scipy.stats.chi2_contingency(table, lambda_="log-likelihood")`` - reports. transCSSR uses the same test. ``"g"`` and ``"chi2"`` use the - chi-squared limit, which is unreliable when expected counts are small. - ``"exact"`` instead compares the G statistic with tables drawn uniformly given - the observed margins whenever an expected count is below 5, and uses the - G-test otherwise. - """ - if test == "tv": - left = counts.state_morph(left_histories) - right = counts.state_morph(right_histories) - distance = 0.5 * sum(abs(left[s] - right[s]) for s in counts.alphabet) - return distance > delta - - table = _contingency_rows(counts, left_histories, right_histories) - if table is None: - return False - return table_significant(table, alpha, test) - - -def morph_test_score( - counts: SuffixCounts, - left_histories: set[History], - right_histories: set[History], - *, - test: MorphTest = "g", -) -> float: - """Score for matching morphs (lower is more similar).""" - if test == "tv": - left = counts.state_morph(left_histories) - right = counts.state_morph(right_histories) - return 0.5 * sum(abs(left[s] - right[s]) for s in counts.alphabet) - table = _contingency_rows(counts, left_histories, right_histories) - if table is None: - return 0.0 - return table_score(table, test) +from sofic.inference.cssr.counts import History, SuffixCounts +from sofic.inference.cssr.significance import ( + MorphTest, + _bonferroni_alpha, + _observed_counts_for_morph, + morph_test_score, + morphs_differ, +) def _cssr_default_lmax(n: int, alphabet_size: int) -> int: @@ -586,172 +413,3 @@ def cssr( homogeneous = _suffix_homogenize(counts, Lmax=max_length, alpha=alpha, test=test, min_count=min_count) return _suffix_reconstruct(homogeneous, counts, seq, Lmax=max_length, alpha=alpha, test=test) - - -#: Significance level used by subtree merging when ``delta = 0`` and to resolve truncated successors. -_SUBTREE_ALPHA = 0.01 - - -def _morph_distance( - counts: SuffixCounts, - left: History, - right: History, - *, - delta: float, -) -> float: - left_morph = counts.morph(left) - right_morph = counts.morph(right) - return 0.5 * sum(abs(left_morph[s] - right_morph[s]) for s in counts.alphabet) - - -def _morphs_equivalent( - counts: SuffixCounts, - left: History, - right: History, - *, - delta: float, - alpha: float = _SUBTREE_ALPHA, - test: MorphTest = "g", -) -> bool: - if delta > 0.0: - return _morph_distance(counts, left, right, delta=delta) <= delta - return not morphs_differ(counts, {left}, {right}, alpha=alpha, test=test) - - -def _cluster_histories_by_morph( - counts: SuffixCounts, - histories: set[History], - *, - delta: float, - alpha: float = _SUBTREE_ALPHA, - test: MorphTest = "g", -) -> dict[int, set[History]]: - parent: dict[History, History] = {history: history for history in histories} - - def find(history: History) -> History: - root = history - while parent[root] != root: - parent[root] = parent[parent[root]] - root = parent[root] - return root - - def union(left: History, right: History) -> None: - left_root = find(left) - right_root = find(right) - if left_root != right_root: - parent[right_root] = left_root - - history_list = sorted(histories) - for index, left in enumerate(history_list): - for right in history_list[index + 1 :]: - if _morphs_equivalent(counts, left, right, delta=delta, alpha=alpha, test=test): - union(left, right) - - clusters: dict[History, set[History]] = defaultdict(set) - for history in histories: - clusters[find(history)].add(history) - - states: dict[int, set[History]] = {} - for state_id, (_root, members) in enumerate(clusters.items()): - states[state_id] = set(members) - return states - - -def subtree_merge( - sequence: Sequence[Any], - *, - L: int | Literal["auto"], - delta: float = 0.0, - alphabet: Sequence[Any] | None = None, - alpha: float = _SUBTREE_ALPHA, - test: MorphTest = "g", - correction: Literal["bonferroni"] | None = None, -) -> EpsilonMachine: - """Reconstruct an ε-machine by merging depth-``L`` subtrees (Crutchfield--Young). - - Histories up to length ``L`` are clustered by next-symbol distribution: within - total-variation distance ``delta``, or, when ``delta = 0``, unless a morph test - (``test``, at level ``alpha``) tells them apart. The clusters are then - determinized as in :func:`cssr`. - - ``L="auto"`` uses :func:`suggest_lmax`. ``correction="bonferroni"`` divides - ``alpha`` by the number of history pairs compared, so that no pair is split - apart by chance; since a rejected test *separates* histories, this makes the - reconstruction more conservative (fewer states). - """ - if L == "auto": - L = suggest_lmax(sequence, alpha=alpha) - if L < 0: - raise ValueError("L must be non-negative") - seq = tuple(sequence) - if len(seq) < 2: - raise ValueError("sequence must contain at least two symbols") - counts = SuffixCounts.from_sequence(seq, alphabet=alphabet, max_length=L + 1) - - histories = {history for history in counts.history_counts if len(history) <= L} - histories.add(()) - if correction == "bonferroni": - alpha /= max(1, len(histories) * (len(histories) - 1) // 2) - elif correction is not None: - raise ValueError(f"unknown correction {correction!r}") - states = list(_cluster_histories_by_morph(counts, histories, delta=delta, alpha=alpha, test=test).values()) - - return _suffix_reconstruct(states, counts, seq, Lmax=L, alpha=alpha, test=test) - - -def spectral( - sequences: Iterable[Any] | None = None, - *, - word_probability: Callable[[Sequence[Any]], float] | None = None, - alphabet: Sequence[Any] | None = None, - rank: int | None = None, - prefix_length: int = 3, - suffix_length: int | None = None, - singular_value_threshold: float = 1e-3, - min_singular_value: float = 1e-12, - max_states: int = 10_000, -) -> EpsilonMachine: - """Reconstruct an ε-machine by spectral learning then mixed-state extraction. - - Learns a weighted finite automaton / observable-operator model from block - statistics :cite:`Balle2014,Hsu2012`, then extracts causal states as the - mixed states of those operators :cite:`Ellison2009`. When the learned - operators are non-negative this is a Mealy projection followed by - :meth:`~sofic.generators.epsilon_machine.EpsilonMachine.from_hmm`; signed - operators use mixed-state enumeration rather than a clustering heuristic. - - Parameters - ---------- - sequences - A single observed realization or an iterable of realizations. Ignored - when ``word_probability`` is given. - word_probability - Optional exact block-probability function ``f(word) -> float``. - ``alphabet`` is then required. - alphabet - Observation alphabet. Inferred from ``sequences`` when omitted. - rank - Number of latent states. When ``None`` the rank is chosen from the - Hankel singular-value spectrum. - prefix_length, suffix_length - Maximum lengths of the prefix and suffix bases. ``suffix_length`` - defaults to ``prefix_length``. - singular_value_threshold, min_singular_value - Cutoffs for automatic rank selection; see - :func:`~sofic.inference.spectral.learn_spectral_wfa`. - max_states - Safety cap on enumerated mixed states. - """ - from sofic.inference.spectral import learn_spectral_wfa, project_to_epsilon_machine - - model = learn_spectral_wfa( - sequences, - word_probability=word_probability, - alphabet=alphabet, - rank=rank, - prefix_length=prefix_length, - suffix_length=suffix_length, - singular_value_threshold=singular_value_threshold, - min_singular_value=min_singular_value, - ) - return project_to_epsilon_machine(model, max_states=max_states) diff --git a/sofic/generators/_morph_tests.py b/sofic/inference/cssr/significance.py similarity index 54% rename from sofic/generators/_morph_tests.py rename to sofic/inference/cssr/significance.py index 518d0a4..398164b 100644 --- a/sofic/generators/_morph_tests.py +++ b/sofic/inference/cssr/significance.py @@ -1,23 +1,31 @@ -"""Significance tests on two-row contingency tables of next-symbol counts. +"""Significance tests that decide whether CSSR histories share a morph. -Shared by process CSSR (:mod:`sofic.generators.epsilon_inference`), stack CSSR, -and transCSSR (:mod:`sofic.generators.epsilon_transducer_inference`). Each table -holds the counts of the symbols following two sets of histories, one row per set. +Two-row contingency tables of next-symbol counts are compared by a G-test, +Pearson chi-squared test, or Monte Carlo exact G-test. Shared by process CSSR +(:mod:`sofic.inference.cssr.process`), stack CSSR (:mod:`sofic.inference.cssr.stack`), +and transCSSR (:mod:`sofic.inference.cssr.transducer`). """ from __future__ import annotations import zlib -from collections.abc import Mapping, Sequence +from collections import Counter +from collections.abc import Callable, Mapping, Sequence from functools import lru_cache from typing import Any, Literal import numpy as np from scipy import stats +from sofic.inference.cssr.counts import History, StateAggregate, SuffixCounts + #: Contingency-table tests: G-test, Pearson chi-squared, or Monte Carlo exact G-test. TableTest = Literal["g", "chi2", "exact"] +#: Morph-equality tests: G-test, chi-squared, total-variation threshold, or Monte Carlo exact G-test. +MorphTest = Literal["g", "chi2", "tv", "exact"] + + #: Monte Carlo tables drawn per ``"exact"`` test. EXACT_DRAWS = 999 @@ -152,3 +160,139 @@ def table_score(table: np.ndarray, test: TableTest) -> float: return float(statistic) statistic = g_statistic(table) return statistic if statistic is not None and np.isfinite(statistic) else 0.0 + + +def _observed_counts_for_morph( + counts: SuffixCounts, + histories: set[History], +) -> Counter[Any]: + observed = Counter() + for history in histories: + observed.update(counts.next_counts.get(history, Counter())) + return observed + + +def _contingency_rows( + counts: SuffixCounts, + left_histories: set[History], + right_histories: set[History], +) -> np.ndarray | None: + return contingency_table( + _observed_counts_for_morph(counts, left_histories), + _observed_counts_for_morph(counts, right_histories), + counts.alphabet, + ) + + +def _bonferroni_alpha( + counts: SuffixCounts, + alpha: float, + *, + max_length: int, + min_count: int, + suffix_length: Callable[[History], int] = len, +) -> float: + """``alpha`` divided by the number of suffixes eligible for a split test.""" + eligible = sum( + 1 + for history, following in counts.next_counts.items() + if 0 < suffix_length(history) <= max_length and sum(following.values()) >= max(1, min_count) + ) + return alpha / max(1, eligible) + + +def morphs_differ( + counts: SuffixCounts, + left_histories: set[History], + right_histories: set[History], + *, + alpha: float = 0.05, + test: MorphTest = "g", + delta: float = 0.0, +) -> bool: + """Return whether two history sets have significantly different morphs. + + ``"g"`` is the G-test (log-likelihood ratio) with Yates' continuity correction + when the table has one degree of freedom, i.e. two observed symbols; this is + the statistic ``scipy.stats.chi2_contingency(table, lambda_="log-likelihood")`` + reports. transCSSR uses the same test. ``"g"`` and ``"chi2"`` use the + chi-squared limit, which is unreliable when expected counts are small. + ``"exact"`` instead compares the G statistic with tables drawn uniformly given + the observed margins whenever an expected count is below 5, and uses the + G-test otherwise. + """ + if test == "tv": + left = counts.state_morph(left_histories) + right = counts.state_morph(right_histories) + distance = 0.5 * sum(abs(left[s] - right[s]) for s in counts.alphabet) + return distance > delta + + table = _contingency_rows(counts, left_histories, right_histories) + if table is None: + return False + return table_significant(table, alpha, test) + + +def morph_test_score( + counts: SuffixCounts, + left_histories: set[History], + right_histories: set[History], + *, + test: MorphTest = "g", +) -> float: + """Score for matching morphs (lower is more similar).""" + if test == "tv": + left = counts.state_morph(left_histories) + right = counts.state_morph(right_histories) + return 0.5 * sum(abs(left[s] - right[s]) for s in counts.alphabet) + table = _contingency_rows(counts, left_histories, right_histories) + if table is None: + return 0.0 + return table_score(table, test) + + +def aggregates_differ( + left: StateAggregate, + right: StateAggregate, + *, + input_alphabet: tuple[Any, ...], + output_alphabet: tuple[Any, ...], + alpha: float, + test: TableTest = "g", +) -> bool: + """Return whether two aggregated morphs differ on ``P(output | ., input)`` for some input. + + Each input symbol's output counts are compared with the same test as process + CSSR (:func:`~sofic.inference.cssr.morphs_differ`); in particular + ``"g"`` is the G-test with Yates' continuity correction at one degree of freedom. + """ + for input_symbol in input_alphabet: + table = contingency_table( + left.get(input_symbol, Counter()), + right.get(input_symbol, Counter()), + output_alphabet, + ) + if table is None: + continue + if table_significant(table, alpha, test): + return True + return False + + +def _aggregate_score( + left: StateAggregate, + right: StateAggregate, + *, + input_alphabet: tuple[Any, ...], + output_alphabet: tuple[Any, ...], +) -> float: + total = 0.0 + for input_symbol in input_alphabet: + table = contingency_table( + left.get(input_symbol, Counter()), + right.get(input_symbol, Counter()), + output_alphabet, + ) + if table is not None: + total += table_score(table, "g") + return total diff --git a/sofic/generators/stack_inference.py b/sofic/inference/cssr/stack.py similarity index 87% rename from sofic/generators/stack_inference.py rename to sofic/inference/cssr/stack.py index 03bd1f6..02299ec 100644 --- a/sofic/generators/stack_inference.py +++ b/sofic/inference/cssr/stack.py @@ -4,95 +4,25 @@ from collections import Counter, defaultdict from collections.abc import Callable, Hashable, Mapping, Sequence -from typing import Any, ClassVar, Literal +from typing import Any, Literal -from sofic.automata.papni import DyckAlphabet, is_well_matched, learn_sofic_dyck_shift_papni +from sofic.automata.learning.papni import DyckAlphabet, is_well_matched, learn_sofic_dyck_shift_papni from sofic.exceptions import StochasticValidationError -from sofic.generators._suffix_counts import infer_alphabet, iter_suffixes -from sofic.generators.epsilon_inference import ( - History, - MorphTest, - SuffixCounts, - _bonferroni_alpha, - _cluster_histories_by_morph, - _cssr_default_lmax, - _recurrent_states, - morph_test_score, - morphs_differ, - suggest_lmax, -) from sofic.generators.stack_hmm import HiddenMarkovStackModel from sofic.graph import ATTR_SYMBOL +from sofic.inference.cssr.counts import ConfigurationHistory, StackSuffixCounts, _push +from sofic.inference.cssr.process import _cssr_default_lmax, _recurrent_states, suggest_lmax +from sofic.inference.cssr.significance import MorphTest, _bonferroni_alpha, morph_test_score, morphs_differ +from sofic.inference.cssr.subtree import _cluster_histories_by_morph from sofic.shifts.sofic_dyck import SoficDyckShift, TransitionRef, transition_ref __all__ = [ - "ConfigurationHistory", - "StackSuffixCounts", "fit_stack_hmm_mle", "learn_stack_hmm_papni", "stack_cssr", "stack_subtree_merge", ] -ConfigurationHistory = tuple[tuple[Any, ...], tuple[Any, ...]] - - -class StackSuffixCounts(SuffixCounts): - """Empirical counts of (suffix, stack) histories and following symbols. - - Shares the morph / comparison machinery of :class:`SuffixCounts`; only the - empty-history key and the sequence-scanning constructor differ. - """ - - empty_history: ClassVar[History] = ((), ()) - - def __init__( - self, - alphabet: tuple[Any, ...], - history_counts: Counter[ConfigurationHistory] | None = None, - next_counts: dict[ConfigurationHistory, Counter[Any]] | None = None, - ) -> None: - super().__init__( - alphabet=alphabet, - history_counts=history_counts if history_counts is not None else Counter(), - next_counts=next_counts if next_counts is not None else defaultdict(Counter), - ) - - @classmethod - def from_sequence( # type: ignore[override] - cls, - sequence: Sequence[Any], - *, - alphabet: DyckAlphabet, - max_length: int | None = None, - max_stack_depth: int = 8, - ) -> StackSuffixCounts: - seq = tuple(sequence) - if not seq: - raise ValueError("sequence must be non-empty") - stacks: list[tuple[Any, ...]] = [] - stack: tuple[Any, ...] = () - for symbol in seq: - if symbol not in alphabet.symbol_alphabet: - raise ValueError(f"symbol {symbol!r} not in Dyck alphabet") - stacks.append(stack) - stack = _push(stack, symbol, alphabet=alphabet, max_stack_depth=max_stack_depth) - counts = cls(alphabet=infer_alphabet(alphabet.symbol_alphabet, None)) - for t, suffix in iter_suffixes(seq, max_length if max_length is not None else len(seq)): - history = (suffix, stacks[t]) - counts.history_counts[history] += 1 - counts.next_counts[history][seq[t]] += 1 - return counts - - -def _push(stack: tuple[Any, ...], symbol: Any, *, alphabet: DyckAlphabet, max_stack_depth: int) -> tuple[Any, ...]: - """The stack after ``symbol``: calls push (dropping the bottom past ``max_stack_depth``), returns pop.""" - if symbol in alphabet.call_alphabet: - return (*(stack[1:] if len(stack) >= max_stack_depth else stack), symbol) - if symbol in alphabet.return_alphabet: - return stack[:-1] - return stack - def _successor_history( history: ConfigurationHistory, @@ -454,11 +384,11 @@ def stack_cssr( ) -> HiddenMarkovStackModel: """Reconstruct a stack HMM via configuration-lifted CSSR. - ``Lmax="auto"`` uses :func:`~sofic.generators.epsilon_inference.suggest_lmax` + ``Lmax="auto"`` uses :func:`~sofic.inference.cssr.suggest_lmax` on the observed symbols. Stack processes generally have infinite Markov order, so treat it as a lower bound on the suffix length the data support. ``test="exact"`` and ``correction="bonferroni"`` are as in - :func:`~sofic.generators.epsilon_inference.cssr`; the correction counts + :func:`~sofic.inference.cssr.cssr`; the correction counts eligible (suffix, stack) configurations. """ seq = tuple(sequence) diff --git a/sofic/inference/cssr/subtree.py b/sofic/inference/cssr/subtree.py new file mode 100644 index 0000000..1e0ee52 --- /dev/null +++ b/sofic/inference/cssr/subtree.py @@ -0,0 +1,125 @@ +"""ε-machine reconstruction by merging depth-``L`` subtrees. + +Follows Crutchfield & Young (PRL 1989; PRE 1994). +""" + +from __future__ import annotations + +from collections import defaultdict +from collections.abc import Sequence +from typing import Any, Literal + +from sofic.generators.epsilon_machine import EpsilonMachine +from sofic.inference.cssr.counts import History, SuffixCounts +from sofic.inference.cssr.process import _suffix_reconstruct, suggest_lmax +from sofic.inference.cssr.significance import MorphTest, morphs_differ + +#: Significance level used by subtree merging when ``delta = 0`` and to resolve truncated successors. +_SUBTREE_ALPHA = 0.01 + + +def _morph_distance( + counts: SuffixCounts, + left: History, + right: History, + *, + delta: float, +) -> float: + left_morph = counts.morph(left) + right_morph = counts.morph(right) + return 0.5 * sum(abs(left_morph[s] - right_morph[s]) for s in counts.alphabet) + + +def _morphs_equivalent( + counts: SuffixCounts, + left: History, + right: History, + *, + delta: float, + alpha: float = _SUBTREE_ALPHA, + test: MorphTest = "g", +) -> bool: + if delta > 0.0: + return _morph_distance(counts, left, right, delta=delta) <= delta + return not morphs_differ(counts, {left}, {right}, alpha=alpha, test=test) + + +def _cluster_histories_by_morph( + counts: SuffixCounts, + histories: set[History], + *, + delta: float, + alpha: float = _SUBTREE_ALPHA, + test: MorphTest = "g", +) -> dict[int, set[History]]: + parent: dict[History, History] = {history: history for history in histories} + + def find(history: History) -> History: + root = history + while parent[root] != root: + parent[root] = parent[parent[root]] + root = parent[root] + return root + + def union(left: History, right: History) -> None: + left_root = find(left) + right_root = find(right) + if left_root != right_root: + parent[right_root] = left_root + + history_list = sorted(histories) + for index, left in enumerate(history_list): + for right in history_list[index + 1 :]: + if _morphs_equivalent(counts, left, right, delta=delta, alpha=alpha, test=test): + union(left, right) + + clusters: dict[History, set[History]] = defaultdict(set) + for history in histories: + clusters[find(history)].add(history) + + states: dict[int, set[History]] = {} + for state_id, (_root, members) in enumerate(clusters.items()): + states[state_id] = set(members) + return states + + +def subtree_merge( + sequence: Sequence[Any], + *, + L: int | Literal["auto"], + delta: float = 0.0, + alphabet: Sequence[Any] | None = None, + alpha: float = _SUBTREE_ALPHA, + test: MorphTest = "g", + correction: Literal["bonferroni"] | None = None, +) -> EpsilonMachine: + """Reconstruct an ε-machine by merging depth-``L`` subtrees (Crutchfield--Young). + + Histories up to length ``L`` are clustered by next-symbol distribution: within + total-variation distance ``delta``, or, when ``delta = 0``, unless a morph test + (``test``, at level ``alpha``) tells them apart. The clusters are then + determinized as in :func:`cssr`. + + ``L="auto"`` uses :func:`suggest_lmax`. ``correction="bonferroni"`` divides + ``alpha`` by the number of history pairs compared, so that no pair is split + apart by chance; since a rejected test *separates* histories, this makes the + reconstruction more conservative (fewer states). + """ + if L == "auto": + L = suggest_lmax(sequence, alpha=alpha) + if L < 0: + raise ValueError("L must be non-negative") + seq = tuple(sequence) + if len(seq) < 2: + raise ValueError("sequence must contain at least two symbols") + counts = SuffixCounts.from_sequence(seq, alphabet=alphabet, max_length=L + 1) + + histories = {history for history in counts.history_counts if len(history) <= L} + histories.add(()) + if correction == "bonferroni": + alpha /= max(1, len(histories) * (len(histories) - 1) // 2) + elif correction is not None: + raise ValueError(f"unknown correction {correction!r}") + states = list(_cluster_histories_by_morph(counts, histories, delta=delta, alpha=alpha, test=test).values()) + + return _suffix_reconstruct(states, counts, seq, Lmax=L, alpha=alpha, test=test) diff --git a/sofic/generators/epsilon_transducer_inference.py b/sofic/inference/cssr/transducer.py similarity index 73% rename from sofic/generators/epsilon_transducer_inference.py rename to sofic/inference/cssr/transducer.py index db779fa..f3cd4cd 100644 --- a/sofic/generators/epsilon_transducer_inference.py +++ b/sofic/inference/cssr/transducer.py @@ -13,139 +13,22 @@ from __future__ import annotations from collections import Counter, defaultdict -from collections.abc import Iterable, Sequence -from dataclasses import dataclass, field +from collections.abc import Sequence from typing import Any, Literal from sofic.exceptions import StochasticValidationError -from sofic.generators._morph_tests import TableTest, contingency_table, table_score, table_significant -from sofic.generators._suffix_counts import infer_alphabet, iter_suffixes -from sofic.generators.epsilon_inference import _recurrent_states, suggest_lmax from sofic.generators.epsilon_transducer import EpsilonTransducer from sofic.graph import ATTR_OUTPUT, ATTR_PROB, ATTR_SYMBOL, TransitionGraph - -JointHistory = tuple[tuple[Any, Any], ...] - - -@dataclass -class JointSuffixCounts: - """Empirical counts of joint pasts and following input-conditioned outputs.""" - - input_alphabet: tuple[Any, ...] - output_alphabet: tuple[Any, ...] - history_counts: Counter[JointHistory] = field(default_factory=Counter) - #: ``next_counts[history][input]`` is a Counter over following output symbols. - next_counts: dict[JointHistory, dict[Any, Counter[Any]]] = field(default_factory=dict) - - @classmethod - def from_sequences( - cls, - inputs: Sequence[Any], - outputs: Sequence[Any], - *, - input_alphabet: Sequence[Any] | None = None, - output_alphabet: Sequence[Any] | None = None, - max_length: int, - ) -> JointSuffixCounts: - xs = tuple(inputs) - ys = tuple(outputs) - if len(xs) != len(ys): - raise ValueError("inputs and outputs must have equal length") - if not xs: - raise ValueError("sequences must be non-empty") - counts = cls( - input_alphabet=infer_alphabet(xs, input_alphabet), output_alphabet=infer_alphabet(ys, output_alphabet) - ) - for t, history in iter_suffixes(tuple(zip(xs, ys, strict=True)), max_length): - counts.history_counts[history] += 1 - counts.next_counts.setdefault(history, {}).setdefault(xs[t], Counter())[ys[t]] += 1 - return counts - - def output_counts(self, histories: set[JointHistory], input_symbol: Any) -> Counter[Any]: - observed: Counter[Any] = Counter() - for history in histories: - by_input = self.next_counts.get(history) - if by_input is None: - continue - counter = by_input.get(input_symbol) - if counter is not None: - observed.update(counter) - return observed - - def state_morph(self, histories: set[JointHistory], input_symbol: Any) -> dict[Any, float]: - """Return ``P(output | histories, input_symbol)``.""" - observed = self.output_counts(histories, input_symbol) - total = sum(observed.values()) - if total == 0: - return {} - return {symbol: observed.get(symbol, 0) / total for symbol in self.output_alphabet} - - -#: Aggregated output counts of a state: ``agg[input_symbol]`` is a Counter over outputs. -StateAggregate = dict[Any, Counter[Any]] - - -def _history_aggregate(counts: JointSuffixCounts, history: JointHistory) -> StateAggregate: - return {input_symbol: Counter(counter) for input_symbol, counter in counts.next_counts.get(history, {}).items()} - - -def _merge_aggregate(target: StateAggregate, source: StateAggregate) -> None: - for input_symbol, counter in source.items(): - target.setdefault(input_symbol, Counter()).update(counter) - - -def aggregates_differ( - left: StateAggregate, - right: StateAggregate, - *, - input_alphabet: tuple[Any, ...], - output_alphabet: tuple[Any, ...], - alpha: float, - test: TableTest = "g", -) -> bool: - """Return whether two aggregated morphs differ on ``P(output | ., input)`` for some input. - - Each input symbol's output counts are compared with the same test as process - CSSR (:func:`~sofic.generators.epsilon_inference.morphs_differ`); in particular - ``"g"`` is the G-test with Yates' continuity correction at one degree of freedom. - """ - for input_symbol in input_alphabet: - table = contingency_table( - left.get(input_symbol, Counter()), - right.get(input_symbol, Counter()), - output_alphabet, - ) - if table is None: - continue - if table_significant(table, alpha, test): - return True - return False - - -def _aggregate_score( - left: StateAggregate, - right: StateAggregate, - *, - input_alphabet: tuple[Any, ...], - output_alphabet: tuple[Any, ...], -) -> float: - total = 0.0 - for input_symbol in input_alphabet: - table = contingency_table( - left.get(input_symbol, Counter()), - right.get(input_symbol, Counter()), - output_alphabet, - ) - if table is not None: - total += table_score(table, "g") - return total - - -def _state_aggregate(counts: JointSuffixCounts, histories: Iterable[JointHistory]) -> StateAggregate: - aggregate: StateAggregate = {} - for history in histories: - _merge_aggregate(aggregate, _history_aggregate(counts, history)) - return aggregate +from sofic.inference.cssr.counts import ( + JointHistory, + JointSuffixCounts, + StateAggregate, + _history_aggregate, + _merge_aggregate, + _state_aggregate, +) +from sofic.inference.cssr.process import _recurrent_states, suggest_lmax +from sofic.inference.cssr.significance import TableTest, _aggregate_score, aggregates_differ def _observed(counts: JointSuffixCounts, history: JointHistory) -> int: @@ -231,7 +114,7 @@ def _edges( Shorter histories move to their one-pair extension; length-``Lmax`` histories drop their oldest pair, and the length-``Lmax + 1`` history is re-tested against the truncated history's state (see - :func:`sofic.generators.epsilon_inference._suffix_edges`). + :func:`sofic.inference.cssr.process._suffix_edges`). """ in_alpha, out_alpha = counts.input_alphabet, counts.output_alphabet history_to_state = {h: index for index in alive for h in states[index]} @@ -415,10 +298,10 @@ def transcssr( states. ``Lmax`` bounds the joint-history depth and ``min_count`` the minimum occurrences before a history is eligible to seed a new state. - ``Lmax="auto"`` applies :func:`~sofic.generators.epsilon_inference.suggest_lmax` + ``Lmax="auto"`` applies :func:`~sofic.inference.cssr.suggest_lmax` to the joint ``(input, output)`` sequence. ``test="exact"`` uses the Monte Carlo exact G-test when expected counts are small (see - :func:`~sofic.generators.epsilon_inference.morphs_differ`), and + :func:`~sofic.inference.cssr.morphs_differ`), and ``correction="bonferroni"`` divides ``alpha`` by the number of (history, input symbol) tests that can split a state. """ diff --git a/sofic/inference/hmm/__init__.py b/sofic/inference/hmm/__init__.py new file mode 100644 index 0000000..04546f4 --- /dev/null +++ b/sofic/inference/hmm/__init__.py @@ -0,0 +1,37 @@ +"""Inference for hidden Markov models. + +Forward/backward/Viterbi decoding plus the Cappe, Moulines & Ryden (2005) +toolbox: fixed-interval smoothing (one- and two-slice marginals), Baum-Welch EM +parameter re-estimation, and the score / observed information via the Fisher +and Louis identities. +""" + +from sofic.inference.hmm.em import baum_welch +from sofic.inference.hmm.filtering import ( + backward, + forward, + log_likelihood, + smooth, + two_slice_marginals, + viterbi, +) +from sofic.inference.hmm.information import ( + free_parameter_labels, + observed_information, + score, + standard_errors, +) + +__all__ = [ + "backward", + "baum_welch", + "forward", + "free_parameter_labels", + "log_likelihood", + "observed_information", + "score", + "smooth", + "standard_errors", + "two_slice_marginals", + "viterbi", +] diff --git a/sofic/inference/hmm/em.py b/sofic/inference/hmm/em.py new file mode 100644 index 0000000..bc2860d --- /dev/null +++ b/sofic/inference/hmm/em.py @@ -0,0 +1,239 @@ +"""Baum-Welch (EM) parameter re-estimation for hidden Markov models. + +Follows Baum, Petrie, Soules & Weiss and Cappe, Moulines & Ryden (2005, Chapter 10). +""" + +from __future__ import annotations + +import warnings +from collections import defaultdict +from collections.abc import Iterable +from typing import Any + +import numpy as np + +from sofic.generators.base import HiddenMarkovModel +from sofic.generators.matrices import emission_tensors +from sofic.inference.hmm.filtering import _backward_scaled, _forward_scaled + + +def _expected_edge_counts( + pi: np.ndarray, + joint: dict[Any, np.ndarray], + obs: list[Any], +) -> tuple[dict[tuple[int, Any, int], float], np.ndarray, np.ndarray, float]: + r"""Expected sufficient statistics for one observation sequence. + + Returns ``(edge_counts, source_totals, gamma0, loglik)`` where + + - ``edge_counts[(i, symbol, j)]`` is + :math:`\sum_t P(X_t = i, Y_t = symbol, X_{t+1} = j \mid Y)`, the expected + number of uses of edge ``i --symbol--> j``; + - ``source_totals[i] = \sum_{t=0}^{n-1} P(X_t = i \mid Y)`` is the expected + number of transitions out of state ``i`` (the Baum-Welch denominator); + - ``gamma0`` is the smoothed marginal of the initial state ``X_0``; + - ``loglik`` is the log-likelihood of the sequence in bits. + + Only edges present in ``joint`` (structural support) receive mass, so the + statistics preserve the model topology. + """ + n_states = len(pi) + alpha_hat, log_scales = _forward_scaled(pi, joint, obs) + if not np.all(np.isfinite(log_scales)): + return {}, np.zeros(n_states), np.zeros(n_states), float("-inf") + beta_hat = _backward_scaled(joint, obs, n_states) + edge_counts: dict[tuple[int, Any, int], float] = {} + source_totals = np.zeros(n_states, dtype=float) + for t, symbol in enumerate(obs): + matrix = joint.get(symbol) + if matrix is None: + continue + block = alpha_hat[t][:, None] * matrix * beta_hat[t + 1][None, :] + total = float(block.sum()) + if total <= 0.0: + continue + block = block / total + source_totals += block.sum(axis=1) + for i, j in np.argwhere(block > 0.0): + key = (int(i), symbol, int(j)) + edge_counts[key] = edge_counts.get(key, 0.0) + float(block[i, j]) + g0 = alpha_hat[0] * beta_hat[0] + s0 = float(g0.sum()) + gamma0 = g0 / s0 if s0 > 0.0 else np.zeros(n_states) + return edge_counts, source_totals, gamma0, float(log_scales.sum()) + + +def _as_sequence_list(sequences: Iterable[Any]) -> list[list[Any]]: + """Normalize ``sequences`` to a list of observation sequences. + + Accepts either a single flat observation sequence (e.g. ``[0, 1, 0]``) or an + iterable of sequences (e.g. ``[[0, 1], [1, 0]]``). A single sequence is + detected when its first element is not itself a non-string sequence. + """ + seqs = list(sequences) + if not seqs: + return [] + first = seqs[0] + if isinstance(first, (list, tuple)) and not isinstance(first, (str, bytes)): + return [list(seq) for seq in seqs] + return [seqs] + + +def baum_welch( + hmm: HiddenMarkovModel, + sequences: Iterable[Any], + *, + max_iter: int = 100, + tol: float = 1e-6, + estimate_initial: bool = True, + n_restarts: int = 1, + rng: np.random.Generator | int | None = None, + return_restarts: bool = False, +) -> tuple[Any, list[float]] | tuple[Any, list[float], list[float]]: + r"""Fit HMM parameters by Baum-Welch (EM) expectation-maximization. + + Re-estimates the Mealy joint edge law + :math:`A_o[i, j] = P(X_{t+1} = j, O = o \mid X_t = i)` and (optionally) the + initial distribution from data, holding the transition-graph topology fixed: + structurally absent edges receive zero expected count and stay absent, so the + fitted model generates the same sofic shift as ``hmm``. This is the EM + algorithm for probabilistic functions of finite Markov chains of Baum, Petrie, + Soules & Weiss and Cappe, Moulines & Ryden (2005, Chapter 10); see also + Rabiner (1989). + + ``sequences`` may be a single observation sequence or an iterable of + sequences (several sequences are needed to identify the initial distribution; + Cappe, Moulines & Ryden, 2005, Section 10.1). Unifilarity is *not* preserved, + so the fit is returned as a plain :class:`~sofic.generators.mealy.MealyHMM`. + + Returns ``(fitted_model, loglik_trace)`` where ``loglik_trace`` is the + non-decreasing sequence of total log-likelihoods (bits) observed before each + parameter update. + + EM converges to a local maximum of the likelihood. With ``n_restarts > 1`` the + first run starts from ``hmm``'s parameters and each further run from edge laws + drawn uniformly (Dirichlet(1)) over each state's structurally allowed edges; + the fit with the highest final log-likelihood is returned. Pass + ``return_restarts=True`` to also get every run's final log-likelihood, which + shows whether near-equal optima exist. + """ + from sofic.generators.mealy import MealyHMM + + mealy = hmm.to_mealy() + idx = mealy.reindex() + n_states = len(idx) + states = [idx.state(i) for i in range(n_states)] + alphabet = frozenset(mealy.observation_alphabet) + seqs = _as_sequence_list(sequences) + + pi, joint = emission_tensors(mealy) + support = { + (i, symbol, j) + for symbol, matrix in joint.items() + for i in range(n_states) + for j in range(n_states) + if matrix[i, j] > 0.0 + } + + if n_restarts < 1: + raise ValueError("n_restarts must be at least 1") + generator = rng if isinstance(rng, np.random.Generator) else np.random.default_rng(rng) + runs = [] + for restart in range(n_restarts): + start_joint = joint if restart == 0 else _random_edge_law(joint, support, n_states, generator) + runs.append( + _baum_welch_run(pi, start_joint, seqs, max_iter=max_iter, tol=tol, estimate_initial=estimate_initial) + ) + finals = [trace[-1] if trace else float("-inf") for _pi, _joint, trace in runs] + pi, joint, loglik_trace = runs[int(np.argmax(finals))] + + fitted = MealyHMM( + initial_distribution={states[i]: float(pi[i]) for i in range(n_states) if pi[i] > 0.0}, + observation_alphabet=alphabet, + ) + for state in states: + fitted.graph.add_state(state) + for i, symbol, j in sorted(support, key=lambda edge: (edge[0], str(edge[1]), edge[2])): + prob = float(joint[symbol][i, j]) + if prob > 0.0: + fitted.add_transition(states[i], states[j], symbol, prob) + fitted.validate() + if return_restarts: + return fitted, loglik_trace, finals + return fitted, loglik_trace + + +def _random_edge_law( + joint: dict[Any, np.ndarray], + support: set[tuple[int, Any, int]], + n_states: int, + rng: np.random.Generator, +) -> dict[Any, np.ndarray]: + """Edge laws drawn uniformly over each state's allowed ``(symbol, target)`` edges.""" + new_joint = {symbol: np.zeros((n_states, n_states), dtype=float) for symbol in joint} + for i in range(n_states): + edges = sorted(((symbol, j) for (source, symbol, j) in support if source == i), key=lambda e: (str(e[0]), e[1])) + if not edges: + continue + weights = rng.dirichlet(np.ones(len(edges))) + for (symbol, j), weight in zip(edges, weights, strict=True): + new_joint[symbol][i, j] = weight + return new_joint + + +def _baum_welch_run( + pi: np.ndarray, + joint: dict[Any, np.ndarray], + seqs: list[Any], + *, + max_iter: int, + tol: float, + estimate_initial: bool, +) -> tuple[np.ndarray, dict[Any, np.ndarray], list[float]]: + """One EM run from ``(pi, joint)``; returns the final parameters and trace.""" + n_states = len(pi) + loglik_trace: list[float] = [] + prev_ll: float | None = None + for _iteration in range(max_iter): + total_edge_counts: dict[tuple[int, Any, int], float] = defaultdict(float) + total_source = np.zeros(n_states, dtype=float) + gamma0_sum = np.zeros(n_states, dtype=float) + total_ll = 0.0 + skipped = 0 + for obs in seqs: + edge_counts, source_totals, gamma0, loglik = _expected_edge_counts(pi, joint, obs) + if not np.isfinite(loglik): + skipped += 1 + continue + for key, value in edge_counts.items(): + total_edge_counts[key] += value + total_source += source_totals + gamma0_sum += gamma0 + total_ll += loglik + if seqs and skipped == len(seqs): + raise ValueError("every observation sequence has zero probability under the model") + if skipped and _iteration == 0: + warnings.warn( + f"{skipped} of {len(seqs)} sequences have zero probability under the model and are ignored", + RuntimeWarning, + stacklevel=3, + ) + loglik_trace.append(total_ll) + if prev_ll is not None and abs(total_ll - prev_ll) < tol: + break + prev_ll = total_ll + + new_joint = {symbol: np.zeros((n_states, n_states), dtype=float) for symbol in joint} + for (i, symbol, j), count in total_edge_counts.items(): + if total_source[i] > 0.0: + new_joint[symbol][i, j] = count / total_source[i] + for i in range(n_states): + if total_source[i] <= 0.0: + for symbol in joint: + new_joint[symbol][i, :] = joint[symbol][i, :] + joint = new_joint + if estimate_initial: + mass = float(gamma0_sum.sum()) + if mass > 0.0: + pi = gamma0_sum / mass + return pi, joint, loglik_trace diff --git a/sofic/inference/hmm/filtering.py b/sofic/inference/hmm/filtering.py new file mode 100644 index 0000000..337b44c --- /dev/null +++ b/sofic/inference/hmm/filtering.py @@ -0,0 +1,237 @@ +"""Forward/backward filtering, smoothing, and Viterbi decoding for hidden Markov models. + +Fixed-interval smoothing (one- and two-slice marginals) follows Cappe, Moulines & +Ryden (2005, Section 3.2). +""" + +from __future__ import annotations + +from collections.abc import Hashable, Sequence +from typing import Any + +import numpy as np + +from sofic.generators.base import HiddenMarkovModel +from sofic.generators.matrices import emission_tensors, symbol_matrices + + +def _as_mealy_hmm(hmm: HiddenMarkovModel) -> Any: + """Return a Mealy-style representation through the HMM representation hook.""" + return hmm.to_mealy() + + +def _forward_scaled( + pi: np.ndarray, + joint: dict[Any, np.ndarray], + obs: list[Any], +) -> tuple[np.ndarray, np.ndarray]: + """Return per-step-normalized forward messages and log scaling factors. + + ``alpha_hat[t]`` sums to one; ``log2 P(obs) = log_scales.sum()``. A ``-inf`` + entry in ``log_scales`` marks an impossible step. Normalizing each step avoids + the underflow that makes the raw forward product vanish for long sequences. + """ + n = len(pi) + alpha_hat = np.zeros((len(obs) + 1, n), dtype=float) + log_scales = np.zeros(len(obs) + 1, dtype=float) + total0 = float(pi.sum()) + if total0 <= 0.0: + log_scales[0] = -np.inf + return alpha_hat, log_scales + alpha_hat[0] = pi / total0 + log_scales[0] = float(np.log2(total0)) + for t, symbol in enumerate(obs): + matrix = joint.get(symbol) + if matrix is None: + log_scales[t + 1] = -np.inf + continue + row = alpha_hat[t] @ matrix + scale = float(row.sum()) + if scale <= 0.0: + log_scales[t + 1] = -np.inf + continue + alpha_hat[t + 1] = row / scale + log_scales[t + 1] = float(np.log2(scale)) + return alpha_hat, log_scales + + +def forward(hmm: HiddenMarkovModel, observations: Sequence[Any], *, scaled: bool = False) -> np.ndarray: + """Return forward messages ``alpha[t, s]`` for ``len(observations)+1`` rows. + + With ``scaled=True`` each row is normalized to sum to one (the numerically + stable message used for posteriors); otherwise the raw messages are returned. + """ + pi, joint = emission_tensors(hmm) + obs = list(observations) + if scaled: + alpha_hat, _log_scales = _forward_scaled(pi, joint, obs) + return alpha_hat + n = len(pi) + alpha = np.zeros((len(obs) + 1, n), dtype=float) + alpha[0] = pi + for t, symbol in enumerate(obs): + matrix = joint.get(symbol) + if matrix is None: + alpha[t + 1] = 0.0 + else: + alpha[t + 1] = alpha[t] @ matrix + return alpha + + +def backward(hmm: HiddenMarkovModel, observations: Sequence[Any], *, scaled: bool = False) -> np.ndarray: + """Return backward messages ``beta[t, s]`` for ``len(observations)+1`` rows. + + With ``scaled=True`` each row is normalized to sum to one. The smoothed + posterior is then ``normalize(alpha_hat[t] * beta_hat[t])`` (the per-row + scaling constants cancel on renormalization). + """ + joint = symbol_matrices(_as_mealy_hmm(hmm)) + n = next(iter(joint.values())).shape[0] if joint else len(_as_mealy_hmm(hmm).reindex()) + obs = list(observations) + beta = np.zeros((len(obs) + 1, n), dtype=float) + beta[len(obs)] = 1.0 + for t in range(len(obs) - 1, -1, -1): + matrix = joint.get(obs[t]) + if matrix is None: + beta[t] = 0.0 + else: + beta[t] = matrix @ beta[t + 1] + if scaled: + total = float(beta[t].sum()) + if total > 0.0: + beta[t] = beta[t] / total + return beta + + +def _backward_scaled(joint: dict[Any, np.ndarray], obs: list[Any], n_states: int) -> np.ndarray: + """Per-row-normalized backward messages from precomputed transition tensors. + + ``beta_hat[t]`` sums to one; the per-row scaling constants cancel against the + forward scaling when the smoothed posterior is renormalized. Shares tensors + with the forward pass so smoothing and EM avoid recomputing them. + """ + beta = np.zeros((len(obs) + 1, n_states), dtype=float) + beta[len(obs)] = 1.0 + for t in range(len(obs) - 1, -1, -1): + matrix = joint.get(obs[t]) + row = beta[t + 1] if matrix is None else matrix @ beta[t + 1] + beta[t] = 0.0 if matrix is None else row + total = float(beta[t].sum()) + if total > 0.0: + beta[t] = beta[t] / total + return beta + + +def log_likelihood(hmm: HiddenMarkovModel, observations: Sequence[Any]) -> float: + """Log-likelihood ``log2 P(observations)`` in bits. + + Uses the per-step-scaled forward recursion so the result stays finite for long + sequences instead of underflowing to ``-inf``. + """ + pi, joint = emission_tensors(hmm) + _alpha_hat, log_scales = _forward_scaled(pi, joint, list(observations)) + if not np.all(np.isfinite(log_scales)): + return float("-inf") + return float(log_scales.sum()) + + +def smooth(hmm: HiddenMarkovModel, observations: Sequence[Any]) -> np.ndarray: + r"""Return fixed-interval smoothed marginals ``gamma[t, s]``. + + ``gamma[t, s] = P(X_t = s \mid Y_{0:n-1})`` for ``t = 0, ..., n`` (there are + ``n + 1`` hidden states behind ``n`` edge emissions). Computed as the + per-row-renormalized product of the scaled forward and backward messages, the + forward-backward smoother of Cappe, Moulines & Ryden (2005, Section 3.2). + Rows for observation sequences of zero probability are returned as zeros. + """ + pi, joint = emission_tensors(hmm) + obs = list(observations) + n_states = len(pi) + alpha_hat, log_scales = _forward_scaled(pi, joint, obs) + if not np.all(np.isfinite(log_scales)): + return np.zeros((len(obs) + 1, n_states), dtype=float) + beta_hat = _backward_scaled(joint, obs, n_states) + gamma = alpha_hat * beta_hat + row_sums = gamma.sum(axis=1, keepdims=True) + with np.errstate(invalid="ignore", divide="ignore"): + gamma = np.where(row_sums > 0.0, gamma / row_sums, 0.0) + return gamma + + +def two_slice_marginals(hmm: HiddenMarkovModel, observations: Sequence[Any]) -> np.ndarray: + r"""Return two-slice smoothed marginals ``xi[t, i, j]``. + + ``xi[t, i, j] = P(X_t = i, X_{t+1} = j \mid Y_{0:n-1})`` for ``t = 0, ..., n-1``, + where the transition at index ``t`` emits ``Y_t`` (Cappe, Moulines & Ryden, + 2005, Section 3.2). Marginalizing over ``j`` recovers ``gamma[t]`` for + ``t < n``. Returns an all-zero tensor for zero-probability sequences. + """ + pi, joint = emission_tensors(hmm) + obs = list(observations) + n_states = len(pi) + xi = np.zeros((len(obs), n_states, n_states), dtype=float) + alpha_hat, log_scales = _forward_scaled(pi, joint, obs) + if not np.all(np.isfinite(log_scales)): + return xi + beta_hat = _backward_scaled(joint, obs, n_states) + for t, symbol in enumerate(obs): + matrix = joint.get(symbol) + if matrix is None: + continue + block = alpha_hat[t][:, None] * matrix * beta_hat[t + 1][None, :] + total = float(block.sum()) + if total > 0.0: + xi[t] = block / total + return xi + + +def _log_probabilities(values: np.ndarray) -> np.ndarray: + log_values = np.full(values.shape, -np.inf, dtype=float) + positive = values > 0.0 + log_values[positive] = np.log(values[positive]) + return log_values + + +def viterbi(hmm: HiddenMarkovModel, observations: Sequence[Any]) -> list[Hashable]: + mealy = _as_mealy_hmm(hmm) + idx = mealy.reindex() + pi, joint = emission_tensors(mealy) + n = len(idx) + obs = list(observations) + if n == 0: + return [] + if not obs: + if not np.any(pi > 0.0): + return [] + return [idx.state(int(np.argmax(pi)))] + + log_pi = _log_probabilities(pi) + viterbi_log = np.full((len(obs), n), -np.inf, dtype=float) + backpointer = np.full((len(obs), n), -1, dtype=int) + + matrix0 = joint.get(obs[0]) + if matrix0 is not None: + log_matrix0 = _log_probabilities(matrix0) + for j in range(n): + best = log_pi + log_matrix0[:, j] + viterbi_log[0, j] = np.max(best) + backpointer[0, j] = int(np.argmax(best)) + + for t in range(1, len(obs)): + matrix = joint.get(obs[t]) + if matrix is None: + continue + log_matrix = _log_probabilities(matrix) + for j in range(n): + scores = viterbi_log[t - 1] + log_matrix[:, j] + viterbi_log[t, j] = np.max(scores) + backpointer[t, j] = int(np.argmax(scores)) + + if not np.any(np.isfinite(viterbi_log[-1])): + return [] + + path = [0] * len(obs) + path[-1] = int(np.argmax(viterbi_log[-1])) + for t in range(len(obs) - 2, -1, -1): + path[t] = backpointer[t + 1, path[t + 1]] + return [idx.state(i) for i in path] diff --git a/sofic/inference/hmm/information.py b/sofic/inference/hmm/information.py new file mode 100644 index 0000000..66b04f9 --- /dev/null +++ b/sofic/inference/hmm/information.py @@ -0,0 +1,208 @@ +"""Score and observed information of hidden Markov models. + +The score follows from the Fisher identity and the observed information from +Louis' identity (Cappe, Moulines & Ryden, 2005, Section 10.2.3). +""" + +from __future__ import annotations + +from collections.abc import Hashable, Sequence +from typing import Any + +import numpy as np + +from sofic.generators.base import HiddenMarkovModel +from sofic.generators.matrices import emission_tensors, symbol_matrices +from sofic.inference.hmm.em import _expected_edge_counts +from sofic.inference.hmm.filtering import _forward_scaled + + +def score(hmm: HiddenMarkovModel, observations: Sequence[Any]) -> dict[tuple[Hashable, Any, Hashable], float]: + r"""Return the score (gradient of the log-likelihood) via the Fisher identity. + + For each edge ``i --o--> j``, returns + :math:`\partial \log P(Y) / \partial A_o[i, j] = E[N_{i,o,j} \mid Y] / A_o[i, j]`, + where ``N`` is the (unobserved) edge-use count. This is Fisher's identity, + ``\nabla \log L(\theta) = E[\nabla \log f(X, Y; \theta) \mid Y]`` (Cappe, + Moulines & Ryden, 2005, Section 10.2.3), evaluated in the raw (unconstrained) + joint-edge parameters. Keys are ``(source, symbol, target)`` state labels. + + The score is the gradient of the *natural* log-likelihood + :math:`\ln P(Y) = \ln 2 \cdot` :func:`log_likelihood`, the convention under + which the observed information and standard errors are standard. + """ + mealy = hmm.to_mealy() + idx = mealy.reindex() + pi, joint = emission_tensors(mealy) + edge_counts, _source_totals, _gamma0, loglik = _expected_edge_counts(pi, joint, list(observations)) + if not np.isfinite(loglik): + raise ValueError("observations have zero probability under the model; score is undefined") + result: dict[tuple[Hashable, Any, Hashable], float] = {} + n_states = len(pi) + for symbol, matrix in joint.items(): + for i in range(n_states): + for j in range(n_states): + prob = float(matrix[i, j]) + if prob > 0.0: + count = edge_counts.get((i, symbol, j), 0.0) + result[(idx.state(i), symbol, idx.state(j))] = count / prob + return result + + +def _free_parameterization( + joint: dict[Any, np.ndarray], + n_states: int, +) -> tuple[list[tuple[int, Any, int]], list[tuple[int, Any, int]], list[int]]: + """Build the free multinomial parameterization of the joint edge law. + + Each source state whose outgoing edges number ``k >= 2`` contributes ``k - 1`` + free parameters (its last edge in canonical order is the reference). Returns + ``(free_edges, reference_by_param, source_by_param)``: the edge for each free + parameter, the reference edge of its source block, and the source-state index. + """ + free_edges: list[tuple[int, Any, int]] = [] + reference_by_param: list[tuple[int, Any, int]] = [] + source_by_param: list[int] = [] + for i in range(n_states): + out_edges = sorted( + ((i, symbol, j) for symbol, matrix in joint.items() for j in range(n_states) if matrix[i, j] > 0.0), + key=lambda edge: (str(edge[1]), edge[2]), + ) + if len(out_edges) < 2: + continue + reference = out_edges[-1] + for edge in out_edges[:-1]: + free_edges.append(edge) + reference_by_param.append(reference) + source_by_param.append(i) + return free_edges, reference_by_param, source_by_param + + +def observed_information(hmm: HiddenMarkovModel, observations: Sequence[Any]) -> np.ndarray: + r"""Return the observed information matrix via Louis' identity. + + The observed information ``J = -\partial^2 \log L / \partial\theta^2`` for the + free multinomial parameters of the joint edge law is obtained from Louis' + (1982) identity, + + .. math:: J = E[-\partial^2 \ell_c \mid Y] - \operatorname{Cov}(\partial \ell_c \mid Y), + + where :math:`\ell_c` is the complete-data log-likelihood (Cappe, Moulines & + Ryden, 2005, Section 10.2.3). The complete-data information ``B`` follows from + the expected edge counts; the conditional covariance of the complete-data + score is computed exactly by a forward smoothing recursion for the first and + second moments of the additive score functional. The matrix is ordered by + :func:`free_parameter_labels`; an empty ``(0, 0)`` matrix is returned when the + model has no free parameters. + """ + mealy = hmm.to_mealy() + pi, joint = emission_tensors(mealy) + obs = list(observations) + n_states = len(pi) + + free_edges, reference_by_param, source_by_param = _free_parameterization(joint, n_states) + d = len(free_edges) + if d == 0: + return np.zeros((0, 0), dtype=float) + + edge_counts, _source_totals, _gamma0, loglik = _expected_edge_counts(pi, joint, obs) + if not np.isfinite(loglik): + raise ValueError("observations have zero probability under the model; information is undefined") + + prob_of = {edge: float(joint[edge[1]][edge[0], edge[2]]) for edge in set(free_edges) | set(reference_by_param)} + + # Complete-data information B = E[-d^2 l_c | Y], block-diagonal by source state. + complete_information = np.zeros((d, d), dtype=float) + for p in range(d): + ref_p = reference_by_param[p] + count_ref = edge_counts.get(ref_p, 0.0) + ref_term = count_ref / prob_of[ref_p] ** 2 + for q in range(d): + if source_by_param[p] != source_by_param[q]: + continue + value = ref_term + if p == q: + edge_p = free_edges[p] + value += edge_counts.get(edge_p, 0.0) / prob_of[edge_p] ** 2 + complete_information[p, q] = value + + # Per-transition score contribution s(edge) as a d-vector (sparse per source block). + edge_score: dict[tuple[int, Any, int], np.ndarray] = {} + for p, edge in enumerate(free_edges): + edge_score.setdefault(edge, np.zeros(d))[p] += 1.0 / prob_of[edge] + for p, ref in enumerate(reference_by_param): + edge_score.setdefault(ref, np.zeros(d))[p] += -1.0 / prob_of[ref] + zero_d = np.zeros(d) + + # Forward smoothing recursion for E[S | Y] and E[S S^T | Y] of the additive + # complete-data score functional S = sum_t s(edge_t). + alpha_hat, _log_scales = _forward_scaled(pi, joint, obs) + first = np.zeros((n_states, d), dtype=float) + second = np.zeros((n_states, d, d), dtype=float) + for t, symbol in enumerate(obs): + matrix = joint.get(symbol) + if matrix is None: + continue + weight = alpha_hat[t][:, None] * matrix # weight[i, k] = P(X_t=i, X_{t+1}=k, Y_t | Y_{0:t-1}) + denom = weight.sum(axis=0) + new_first = np.zeros((n_states, d), dtype=float) + new_second = np.zeros((n_states, d, d), dtype=float) + for k in range(n_states): + if denom[k] <= 0.0: + continue + for i in range(n_states): + if weight[i, k] <= 0.0: + continue + retro = weight[i, k] / denom[k] # P(X_t=i | X_{t+1}=k, Y_{0:t}) + s_vec = edge_score.get((i, symbol, k), zero_d) + first_i = first[i] + combined = first_i + s_vec + new_first[k] += retro * combined + cross = np.outer(first_i, s_vec) + new_second[k] += retro * (second[i] + cross + cross.T + np.outer(s_vec, s_vec)) + first, second = new_first, new_second + + phi_final = alpha_hat[len(obs)] + expected_score = phi_final @ first + expected_outer = np.einsum("k,kpq->pq", phi_final, second) + score_covariance = expected_outer - np.outer(expected_score, expected_score) + return complete_information - score_covariance + + +def free_parameter_labels(hmm: HiddenMarkovModel) -> list[tuple[Hashable, Any, Hashable]]: + """Return the ``(source, symbol, target)`` label for each free parameter. + + The order matches the rows and columns of :func:`observed_information` and the + entries of :func:`standard_errors`. + """ + mealy = hmm.to_mealy() + idx = mealy.reindex() + joint = symbol_matrices(mealy) + free_edges, _reference, _source = _free_parameterization(joint, len(idx)) + return [(idx.state(i), symbol, idx.state(j)) for i, symbol, j in free_edges] + + +def standard_errors( + hmm: HiddenMarkovModel, + observations: Sequence[Any], +) -> dict[tuple[Hashable, Any, Hashable], float]: + r"""Return asymptotic standard errors of the free edge parameters. + + Standard errors are ``sqrt(diag(J^{-1}))`` where ``J`` is the + :func:`observed_information` matrix (Cappe, Moulines & Ryden, 2005, + Section 10.2.3). Uses the Moore-Penrose pseudoinverse when ``J`` is singular; + a non-positive variance estimate (numerically unidentified parameter) yields + ``nan``. Keyed by the labels from :func:`free_parameter_labels`. + """ + labels = free_parameter_labels(hmm) + information = observed_information(hmm, observations) + if information.shape[0] == 0: + return {} + try: + covariance = np.linalg.inv(information) + except np.linalg.LinAlgError: + covariance = np.linalg.pinv(information) + variances = np.diag(covariance) + with np.errstate(invalid="ignore"): + errors = np.where(variances > 0.0, np.sqrt(variances), np.nan) + return dict(zip(labels, (float(value) for value in errors), strict=True)) diff --git a/sofic/inference/model_selection.py b/sofic/inference/model_selection.py index cc4345c..d48dc9a 100644 --- a/sofic/inference/model_selection.py +++ b/sofic/inference/model_selection.py @@ -8,7 +8,7 @@ evidences of :mod:`sofic.inference.bayesian`: they score any fitted :class:`~sofic.generators.base.HiddenMarkovModel` (ε-machine, Mealy HMM, Markov chain) using the log-likelihood (in bits) from -:func:`sofic.generators.hmm_inference.log_likelihood` and a free-parameter count +:func:`sofic.inference.hmm.log_likelihood` and a free-parameter count read off the transition graph, so they are likelihood-agnostic and apply directly to discrete-emission models. @@ -30,7 +30,7 @@ import numpy as np from sofic.generators.base import HiddenMarkovModel -from sofic.generators.hmm_inference import free_parameter_labels, log_likelihood +from sofic.inference.hmm import free_parameter_labels, log_likelihood __all__ = [ "ModelScores", @@ -94,7 +94,7 @@ def count_free_parameters(model: HiddenMarkovModel, *, include_initial: bool = F Counts one free parameter per non-reference outgoing edge at each state (the multinomial free-parameterization of the joint emission-transition law used - by :func:`sofic.generators.hmm_inference.observed_information`). With + by :func:`sofic.inference.hmm.observed_information`). With ``include_initial`` the ``n - 1`` free parameters of the initial distribution are added; for a stationary presentation the initial law is determined by the dynamics, so this defaults to ``False``. diff --git a/sofic/inference/spectral.py b/sofic/inference/spectral.py index abf9b23..beccbd7 100644 --- a/sofic/inference/spectral.py +++ b/sofic/inference/spectral.py @@ -33,12 +33,15 @@ from collections import defaultdict from collections.abc import Callable, Hashable, Iterable, Sequence -from typing import Any +from typing import TYPE_CHECKING, Any import numpy as np from sofic.generators.quasi_realization import QuasiRealization +if TYPE_CHECKING: + from sofic.generators.epsilon_machine import EpsilonMachine + __all__ = [ "SpectralInferenceError", "hankel_matrices", @@ -46,6 +49,7 @@ "project_to_epsilon_machine", "project_to_mealy", "project_to_nmachine", + "spectral", "spectral_singular_values", ] @@ -62,7 +66,7 @@ def _normalize_sequences(sequences: Iterable[Any]) -> list[tuple[Any, ...]]: Accepts either a single flat observation sequence (e.g. ``[0, 1, 0]``) or an iterable of sequences (e.g. ``[[0, 1], [1, 0]]``), mirroring - :func:`sofic.generators.hmm_inference.baum_welch`. + :func:`sofic.inference.hmm.baum_welch`. """ seqs = list(sequences) if not seqs: @@ -562,3 +566,59 @@ def register(state: MixedState) -> MixedState: initial_distribution=initial, observation_alphabet=frozenset(symbols), ) + + +def spectral( + sequences: Iterable[Any] | None = None, + *, + word_probability: Callable[[Sequence[Any]], float] | None = None, + alphabet: Sequence[Any] | None = None, + rank: int | None = None, + prefix_length: int = 3, + suffix_length: int | None = None, + singular_value_threshold: float = 1e-3, + min_singular_value: float = 1e-12, + max_states: int = 10_000, +) -> EpsilonMachine: + """Reconstruct an ε-machine by spectral learning then mixed-state extraction. + + Learns a weighted finite automaton / observable-operator model from block + statistics :cite:`Balle2014,Hsu2012`, then extracts causal states as the + mixed states of those operators :cite:`Ellison2009`. When the learned + operators are non-negative this is a Mealy projection followed by + :meth:`~sofic.generators.epsilon_machine.EpsilonMachine.from_hmm`; signed + operators use mixed-state enumeration rather than a clustering heuristic. + + Parameters + ---------- + sequences + A single observed realization or an iterable of realizations. Ignored + when ``word_probability`` is given. + word_probability + Optional exact block-probability function ``f(word) -> float``. + ``alphabet`` is then required. + alphabet + Observation alphabet. Inferred from ``sequences`` when omitted. + rank + Number of latent states. When ``None`` the rank is chosen from the + Hankel singular-value spectrum. + prefix_length, suffix_length + Maximum lengths of the prefix and suffix bases. ``suffix_length`` + defaults to ``prefix_length``. + singular_value_threshold, min_singular_value + Cutoffs for automatic rank selection; see + :func:`~sofic.inference.spectral.learn_spectral_wfa`. + max_states + Safety cap on enumerated mixed states. + """ + model = learn_spectral_wfa( + sequences, + word_probability=word_probability, + alphabet=alphabet, + rank=rank, + prefix_length=prefix_length, + suffix_length=suffix_length, + singular_value_threshold=singular_value_threshold, + min_singular_value=min_singular_value, + ) + return project_to_epsilon_machine(model, max_states=max_states) diff --git a/sofic/serialization.py b/sofic/serialization.py index c660409..a01cdca 100644 --- a/sofic/serialization.py +++ b/sofic/serialization.py @@ -298,12 +298,12 @@ def _spec(cls: type[StateMachine], fields: tuple[str, ...], builder: str = "defa @cache def _specs() -> tuple[_ModelSpec, ...]: - from sofic.automata.atomaton import Atomaton, AtomicAutomaton, MaximizedPrimeAtomaton from sofic.automata.buchi import BuchiAutomaton + from sofic.automata.canonical.atomaton import Atomaton, AtomicAutomaton, MaximizedPrimeAtomaton + from sofic.automata.canonical.rfsa import CanonicalRFSA, ResidualFiniteStateAutomaton from sofic.automata.dfa import DFA from sofic.automata.nfa import NFA from sofic.automata.nwa import NestedWordAutomaton - from sofic.automata.rfsa import CanonicalRFSA, ResidualFiniteStateAutomaton from sofic.automata.subsequential import SubsequentialTransducer, WeightedFiniteStateTransducer from sofic.automata.transducers import MealyMachine, MooreMachine from sofic.automata.unifilar import UnifilarAutomaton diff --git a/sofic/testing/strategies.py b/sofic/testing/strategies.py index cab8e14..d84efb3 100644 --- a/sofic/testing/strategies.py +++ b/sofic/testing/strategies.py @@ -11,7 +11,7 @@ from typing import Any from sofic.automata.dfa import DFA -from sofic.automata.icdfa import icdfa_string_to_dfa, iter_icdfa_empty_strings +from sofic.automata.enumeration.icdfa import icdfa_string_to_dfa, iter_icdfa_empty_strings from sofic.generators.epsilon_machine import EpsilonMachine from sofic.generators.topological_epsilon_enumeration import ( idfa_string_to_epsilon_machine, diff --git a/tests/test_active_learning.py b/tests/test_active_learning.py index 0d58d1e..6bdf3d6 100644 --- a/tests/test_active_learning.py +++ b/tests/test_active_learning.py @@ -6,7 +6,8 @@ import pytest -from sofic.automata.active import ( +from sofic.automata.dfa import DFA +from sofic.automata.learning.active import ( ExhaustiveEquivalenceOracle, FunctionMembershipOracle, LanguageMembershipOracle, @@ -16,7 +17,6 @@ learn_dfa_ttt, learn_mealy_from_transducer, ) -from sofic.automata.dfa import DFA from sofic.automata.transducers import MealyMachine ALPHABET = ("a", "b") @@ -159,7 +159,7 @@ def test_mealy_learns_last_symbol_echo(): def echo(word): return tuple(word) - from sofic.automata.active import FunctionMealyOracle, MealyExhaustiveEquivalenceOracle, learn_mealy_lstar + from sofic.automata.learning.active import FunctionMealyOracle, MealyExhaustiveEquivalenceOracle, learn_mealy_lstar oracle = FunctionMealyOracle(echo) alphabet = ("0", "1") diff --git a/tests/test_atomaton.py b/tests/test_atomaton.py index eacd35d..a720581 100644 --- a/tests/test_atomaton.py +++ b/tests/test_atomaton.py @@ -2,10 +2,10 @@ import pytest -from sofic.automata.atomaton import Atomaton, MaximizedPrimeAtomaton, atomic_states, is_atomic +from sofic.automata.canonical.atomaton import Atomaton, MaximizedPrimeAtomaton, atomic_states, is_atomic +from sofic.automata.canonical.rfsa import CanonicalRFSA from sofic.automata.dfa import DFA from sofic.automata.nfa import NFA -from sofic.automata.rfsa import CanonicalRFSA from sofic.exceptions import SoficValidationError diff --git a/tests/test_canonical_rfsa.py b/tests/test_canonical_rfsa.py index 8f167b9..6143a70 100644 --- a/tests/test_canonical_rfsa.py +++ b/tests/test_canonical_rfsa.py @@ -4,13 +4,13 @@ from hypothesis import given, settings from sofic.automata.algorithms import equivalent -from sofic.automata.atomaton import Atomaton, MaximizedPrimeAtomaton, is_atomic +from sofic.automata.canonical.atomaton import Atomaton, MaximizedPrimeAtomaton, is_atomic +from sofic.automata.canonical.rfsa import CanonicalRFSA, ResidualFiniteStateAutomaton from sofic.automata.dfa import DFA from sofic.automata.languages.atoms import atoms, prime_atoms from sofic.automata.languages.base import AutomatonLanguage from sofic.automata.languages.residuals import prime_residuals from sofic.automata.nfa import NFA -from sofic.automata.rfsa import CanonicalRFSA, ResidualFiniteStateAutomaton from sofic.exceptions import SoficValidationError from sofic.testing.strategies import dfas @@ -90,7 +90,7 @@ def test_canonical_rfsa_and_prime_atomaton_recognize_the_language(dfa): def test_maximized_prime_atomaton_need_not_be_atomic(): """Its right languages lie between an atom and a maximized atom (Tamm 2015), not on atoms.""" - from sofic.automata.icdfa import icdfa_string_to_dfa + from sofic.automata.enumeration.icdfa import icdfa_string_to_dfa dfa = icdfa_string_to_dfa((0, 1, 0, 2, 0, 1), ("0", "1"), n=3, k=2, final_states=frozenset({0, 1})) mpa = MaximizedPrimeAtomaton.from_language(dfa) diff --git a/tests/test_epsilon_inference.py b/tests/test_epsilon_inference.py index 532b49d..0519367 100644 --- a/tests/test_epsilon_inference.py +++ b/tests/test_epsilon_inference.py @@ -10,10 +10,11 @@ import pytest from sofic.examples.epsilon_machines import bernoulli, even_process, golden_mean -from sofic.generators.epsilon_inference import cssr, spectral, subtree_merge from sofic.generators.epsilon_machine import EpsilonMachine -from sofic.generators.hmm_inference import sample +from sofic.generators.sampling import sample from sofic.graph import ATTR_EMISSION, ATTR_PROB +from sofic.inference.cssr import cssr, subtree_merge +from sofic.inference.spectral import spectral def _transition_signature(hmm: EpsilonMachine) -> dict[Hashable, tuple[tuple[Any, Hashable, float], ...]]: @@ -289,7 +290,7 @@ def _has_markov_order_selection() -> bool: @needs_markov_order def test_suggest_lmax_markov_sources(): from sofic.examples import processes - from sofic.generators.epsilon_inference import suggest_lmax + from sofic.inference.cssr import suggest_lmax observations, _ = sample(golden_mean(0.5), 4000, np.random.default_rng(1)) assert suggest_lmax(observations) == 1 @@ -300,7 +301,7 @@ def test_suggest_lmax_markov_sources(): @needs_markov_order def test_suggest_lmax_grows_for_even_process(): """The even process has infinite Markov order, so the suggestion grows with data.""" - from sofic.generators.epsilon_inference import suggest_lmax + from sofic.inference.cssr import suggest_lmax short, _ = sample(even_process(0.5), 300, np.random.default_rng(3)) long, _ = sample(even_process(0.5), 30000, np.random.default_rng(3)) @@ -317,7 +318,7 @@ def test_cssr_auto_lmax_golden_mean(rng: np.random.Generator): def test_exact_morph_test_small_counts(): """With tiny counts the exact test is calibrated where the chi-squared limit is not.""" - from sofic.generators.epsilon_inference import SuffixCounts, morphs_differ + from sofic.inference.cssr import SuffixCounts, morphs_differ rng = np.random.default_rng(4) rejections = {"g": 0, "exact": 0} diff --git a/tests/test_epsilon_transducer_inference.py b/tests/test_epsilon_transducer_inference.py index be5de34..9c331ba 100644 --- a/tests/test_epsilon_transducer_inference.py +++ b/tests/test_epsilon_transducer_inference.py @@ -9,7 +9,7 @@ from sofic import EpsilonTransducer, MealyHMM from sofic.automata.transducer_operations import compose_tg from sofic.examples.processes import BinaryChannel, Delay -from sofic.generators.epsilon_transducer_inference import JointSuffixCounts, transcssr +from sofic.inference.cssr import JointSuffixCounts, transcssr def _iid_input() -> MealyHMM: @@ -159,7 +159,7 @@ def test_shared_g_statistic_matches_scipy_log_likelihood(table): """The G-test shared with process CSSR is scipy's log-likelihood statistic, Yates-corrected at dof 1.""" from scipy import stats - from sofic.generators._morph_tests import g_statistic + from sofic.inference.cssr.significance import g_statistic expected, _p, _dof, _ = stats.chi2_contingency(table, lambda_="log-likelihood") assert g_statistic(table.astype(float)) == pytest.approx(expected, rel=1e-9, abs=1e-12) diff --git a/tests/test_hmm_inference.py b/tests/test_hmm_inference.py index 99403a6..7260272 100644 --- a/tests/test_hmm_inference.py +++ b/tests/test_hmm_inference.py @@ -6,24 +6,24 @@ import pytest from sofic.examples import fair_coin, golden_mean -from sofic.generators.hmm_inference import ( - _forward_scaled, +from sofic.generators.matrices import emission_tensors +from sofic.generators.mealy import MealyHMM +from sofic.generators.sampling import sample +from sofic.graph import ATTR_EMISSION, ATTR_PROB +from sofic.inference.hmm import ( backward, baum_welch, forward, free_parameter_labels, log_likelihood, observed_information, - sample, score, smooth, standard_errors, two_slice_marginals, viterbi, ) -from sofic.generators.matrices import emission_tensors -from sofic.generators.mealy import MealyHMM -from sofic.graph import ATTR_EMISSION, ATTR_PROB +from sofic.inference.hmm.filtering import _forward_scaled def test_forward_coin_initial_and_likelihood(): diff --git a/tests/test_icdfa.py b/tests/test_icdfa.py index 87ed1b8..edb1f22 100644 --- a/tests/test_icdfa.py +++ b/tests/test_icdfa.py @@ -7,7 +7,7 @@ import pytest from sofic.automata.dfa import DFA -from sofic.automata.icdfa import ( +from sofic.automata.enumeration.icdfa import ( ICDFAEnumerationError, _upper_bound_at, count_flag_sequences, diff --git a/tests/test_inference_diagnostics.py b/tests/test_inference_diagnostics.py index f87c31b..fbc5ab3 100644 --- a/tests/test_inference_diagnostics.py +++ b/tests/test_inference_diagnostics.py @@ -6,8 +6,8 @@ import pytest from sofic.examples.epsilon_machines import even_process, golden_mean -from sofic.generators.epsilon_inference import cssr -from sofic.generators.hmm_inference import sample +from sofic.generators.sampling import sample +from sofic.inference.cssr import cssr from sofic.inference.diagnostics import ( goodness_of_fit, reconstruction_sweep, diff --git a/tests/test_learning.py b/tests/test_learning.py index 98d56bd..6613153 100644 --- a/tests/test_learning.py +++ b/tests/test_learning.py @@ -3,13 +3,13 @@ import pytest from hypothesis import given, settings -from sofic.automata.active import AutomatonEquivalenceOracle, LanguageMembershipOracle from sofic.automata.algorithms import equivalent -from sofic.automata.atomaton import MaximizedPrimeAtomaton +from sofic.automata.canonical.atomaton import MaximizedPrimeAtomaton +from sofic.automata.canonical.rfsa import CanonicalRFSA from sofic.automata.dfa import DFA -from sofic.automata.learning import learn_prime_atomaton_nlstar, learn_rfsa_from_language, learn_rfsa_nlstar +from sofic.automata.learning.active import AutomatonEquivalenceOracle, LanguageMembershipOracle +from sofic.automata.learning.nlstar import learn_prime_atomaton_nlstar, learn_rfsa_from_language, learn_rfsa_nlstar from sofic.automata.nfa import NFA -from sofic.automata.rfsa import CanonicalRFSA from sofic.testing.strategies import dfas diff --git a/tests/test_observation.py b/tests/test_observation.py index 5d1753b..b062b2c 100644 --- a/tests/test_observation.py +++ b/tests/test_observation.py @@ -1,7 +1,7 @@ """Tests for observation tables.""" from sofic.automata.languages.base import AutomatonLanguage -from sofic.automata.observation import ObservationTable +from sofic.automata.learning.observation import ObservationTable def test_defaults(): diff --git a/tests/test_rfsa.py b/tests/test_rfsa.py index b0c3cde..dfead3d 100644 --- a/tests/test_rfsa.py +++ b/tests/test_rfsa.py @@ -1,8 +1,8 @@ """Tests for residual and canonical RFSA skeletons.""" +from sofic.automata.canonical.rfsa import CanonicalRFSA, ResidualFiniteStateAutomaton from sofic.automata.dfa import DFA -from sofic.automata.observation import ObservationTable -from sofic.automata.rfsa import CanonicalRFSA, ResidualFiniteStateAutomaton +from sofic.automata.learning.observation import ObservationTable def test_rfsa_validate(): diff --git a/tests/test_stack_inference.py b/tests/test_stack_inference.py index d2290e6..037cc2d 100644 --- a/tests/test_stack_inference.py +++ b/tests/test_stack_inference.py @@ -5,21 +5,21 @@ import numpy as np import pytest -from sofic.automata.papni import DyckAlphabet, is_well_matched, learn_sofic_dyck_shift_papni, papni_encode -from sofic.automata.rpni import learn_dfa_rpni +from sofic.automata.learning.papni import DyckAlphabet, is_well_matched, learn_sofic_dyck_shift_papni, papni_encode +from sofic.automata.learning.rpni import learn_dfa_rpni from sofic.examples.shifts import dyck_shift_order, motzkin_shift -from sofic.generators.epsilon_inference import cssr from sofic.generators.stack_hmm import HiddenMarkovStackModel -from sofic.generators.stack_inference import ( +from sofic.inference.bayesian.stack_hmm import ( + ModelComparisonStackHMM, + StackHMMPosterior, +) +from sofic.inference.cssr import cssr +from sofic.inference.cssr.stack import ( fit_stack_hmm_mle, learn_stack_hmm_papni, stack_cssr, stack_subtree_merge, ) -from sofic.inference.bayesian.stack_hmm import ( - ModelComparisonStackHMM, - StackHMMPosterior, -) from sofic.shifts.dyck_enumeration import ( count_dyck_graph_strings, dyck_graph_string_to_shift, diff --git a/tests/test_testing_strategies.py b/tests/test_testing_strategies.py index 0db589e..a55e429 100644 --- a/tests/test_testing_strategies.py +++ b/tests/test_testing_strategies.py @@ -5,7 +5,7 @@ import pytest from hypothesis import given, settings -from sofic.automata.icdfa import dfa_to_icdfa_string +from sofic.automata.enumeration.icdfa import dfa_to_icdfa_string from sofic.testing.strategies import dfas, epsilon_machines diff --git a/tests/test_topological_epsilon.py b/tests/test_topological_epsilon.py index 840f1df..1725cf6 100644 --- a/tests/test_topological_epsilon.py +++ b/tests/test_topological_epsilon.py @@ -4,7 +4,7 @@ import pytest -from sofic.automata.idfa import ( +from sofic.automata.enumeration.idfa import ( MISSING_TRANSITION, count_accessible_idfa, first_idfa_string, @@ -125,6 +125,6 @@ def test_accessible_idfa_count_grows() -> None: def count_icdfa_placeholder(k: int, n: int) -> int: - from sofic.automata.icdfa import count_icdfa_empty + from sofic.automata.enumeration.icdfa import count_icdfa_empty return count_icdfa_empty(k, n) diff --git a/tests/test_vpa_constructions.py b/tests/test_vpa_constructions.py index cb34e65..cf150fe 100644 --- a/tests/test_vpa_constructions.py +++ b/tests/test_vpa_constructions.py @@ -14,7 +14,7 @@ SingleEntryVisiblyPushdownAutomaton, VisiblyPushdownAutomaton, ) -from sofic.automata.vpa_constructions import to_multiple_entry, to_single_entry +from sofic.automata.vpa.operations import to_multiple_entry, to_single_entry from sofic.exceptions import NonWellMatchedLanguageError CALLS, RETURNS, INTERNALS = frozenset({"c"}), frozenset({"r"}), frozenset({"i"}) diff --git a/tests/test_yaml.py b/tests/test_yaml.py index 0ce8009..3dceaed 100644 --- a/tests/test_yaml.py +++ b/tests/test_yaml.py @@ -5,7 +5,7 @@ import numpy as np import pytest -from sofic.automata.atomaton import Atomaton +from sofic.automata.canonical.atomaton import Atomaton from sofic.automata.dfa import DFA from sofic.automata.nfa import NFA from sofic.automata.nwa import NestedWordAutomaton