From 15bba614b7151c45c3ea571d6ec362eb275ca29b Mon Sep 17 00:00:00 2001 From: Ryan James Date: Sun, 27 Sep 2026 19:16:19 -0600 Subject: [PATCH 1/2] Fix CSSR reconstruction, seeded sampling, default initial distributions, and is_symbolic cost - cssr: grow suffixes into the past, keep suffixes of every length, remove transient states before determinizing, and resolve truncated successors with the length-Lmax+1 suffix (Shalizi et al. 2002). Always returns a valid recurrent machine; conservative default Lmax; min_count now skips rare suffixes; default alpha 0.05 -> 0.01. - Emission tensors use a repr-sorted symbol order, so seeded sampling no longer depends on PYTHONHASHSEED. - A model without an initial distribution starts from its stationary distribution; MealyHMM.add_transition grows the observation alphabet. - is_symbolic/as_prob check sys.modules instead of retrying a failed sympy import on every call (tetris_bag entropy_rate 4.5 s -> 0.04 s). Co-authored-by: Cursor --- docs/generators/epsilon_inference.rst | 23 +- sofic/generators/epsilon_inference.py | 306 ++++++++++++++++++++++++-- sofic/generators/hmm_inference.py | 19 +- sofic/generators/mealy.py | 5 +- sofic/generators/prob.py | 15 +- tests/test_epsilon_inference.py | 55 +++++ tests/test_hmm_inference.py | 25 +++ tests/test_mealy_hmm.py | 16 ++ tests/test_prob.py | 25 +++ 9 files changed, 452 insertions(+), 37 deletions(-) create mode 100644 tests/test_prob.py diff --git a/docs/generators/epsilon_inference.rst b/docs/generators/epsilon_inference.rst index 8586acd..6a12d26 100644 --- a/docs/generators/epsilon_inference.rst +++ b/docs/generators/epsilon_inference.rst @@ -39,10 +39,25 @@ CSSR CSSR :cite:`Shalizi2002` starts from an IID model and grows causal states in three phases: 1. **Initialize** — one state for the empty history. -2. **Homogenize** — extend suffixes; split or assign child histories when next-symbol - distributions differ (G-test, :math:`\chi^2`, or total-variation threshold). -3. **Determinize** — split homogeneous states until transitions are unifilar; drop - transient bottom-SCC states. +2. **Homogenize** — extend each suffix one symbol into the past, up to ``Lmax``; a + child suffix whose next-symbol distribution differs significantly from its + state's (G-test, :math:`\chi^2`, or total-variation threshold) moves to the best + matching state, or starts a new one. States keep suffixes of every length. +3. **Determinize** — drop transient states, then split states until each state and + symbol lead to a single successor, then keep the most-visited recurrent class. + +A length-``Lmax`` suffix has no one-symbol extension in the suffix tree, so its +successor drops the oldest symbol. For a non-Markovian process that can forget the +phase: in the even process with ``Lmax = 3``, the successor of ``011`` on ``1`` would +be the ambiguous ``111``. So the length-``Lmax + 1`` suffix (here ``0111``) is tested +against the truncated suffix's state, and is sent to the best matching state when +the two differ. + +Choose ``Lmax`` at least the synchronization length of the source (its order, for +a Markov source). Much larger values run many more significance tests, and some +split states by chance; lowering ``alpha`` counters this. A process that is not +exactly synchronizable has no finite-``Lmax`` reconstruction, and CSSR returns +extra states. .. autofunction:: cssr diff --git a/sofic/generators/epsilon_inference.py b/sofic/generators/epsilon_inference.py index 4cc97f3..7a7caf0 100644 --- a/sofic/generators/epsilon_inference.py +++ b/sofic/generators/epsilon_inference.py @@ -510,37 +510,309 @@ def _counts_to_mealy( return machine +def _cssr_default_lmax(n: int, alphabet_size: int) -> int: + """A third of ``log_k n``, between 1 and 10. + + Each length-``L`` word is then seen about ``n ** (2/3)`` times. Longer suffixes + multiply the number of significance tests, and with them the false splits. + """ + k = max(2, alphabet_size) + return max(1, min(10, int(np.log(n) / (3 * np.log(k))))) + + +def _suffix_homogenize( + counts: SuffixCounts, + *, + Lmax: int, + alpha: float, + test: Literal["g", "chi2", "tv"], + min_count: int = 1, +) -> list[set[History]]: + """CSSR homogenization: grow suffixes one symbol into the past, up to length ``Lmax``. + + Each child suffix ``a x`` stays in its parent's state unless its next-symbol + distribution differs significantly; then it joins the most similar state that + does not differ, or starts a new one. As in :cite:`Shalizi2002`, states keep + the suffixes of every length they collect. Suffixes seen fewer than ``min_count`` + times are not tested: the significance test is unreliable on so few counts. + """ + states: list[set[History]] = [{()}] + for length in range(Lmax): + for parent_id in range(len(states)): + parent = states[parent_id] + for history in sorted((h for h in parent if len(h) == length), key=repr): + for symbol in counts.alphabet: + child = (symbol, *history) + if sum(counts.next_counts.get(child, Counter()).values()) < max(1, min_count): + continue + if not morphs_differ(counts, parent, {child}, alpha=alpha, test=test): + parent.add(child) + continue + best_id, best_score = None, float("inf") + for candidate_id, candidate in enumerate(states): + if candidate_id == parent_id: + continue + if morphs_differ(counts, candidate, {child}, alpha=alpha, test=test): + continue + score = morph_test_score(counts, candidate, {child}, test=test) + if score < best_score: + best_id, best_score = candidate_id, score + if best_id is None: + states.append({child}) + else: + states[best_id].add(child) + return states + + +def _suffix_successor(history: History, symbol: Any, Lmax: int) -> History: + """The suffix that follows ``history`` on ``symbol``: extended, or truncated at ``Lmax``.""" + extended = (*history, symbol) + return extended[1:] if len(extended) > Lmax else extended + + +def _suffix_edges( + states: list[set[History]], + counts: SuffixCounts, + alive: set[int], + *, + Lmax: int, + alpha: float, + test: Literal["g", "chi2", "tv"], + resolve: bool = True, +) -> dict[int, dict[Any, dict[int, set[History]]]]: + """Successor states of each alive state, by symbol, with the suffixes that lead there. + + A suffix shorter than ``Lmax`` moves to the state holding its one-symbol extension. + A length-``Lmax`` suffix must drop its oldest symbol, which can forget the phase + of a non-Markovian process: for the even process, the truncation of ``0111`` is + ``111``, whose parity is unknown. So the length-``Lmax + 1`` suffix is tested + against the truncated suffix's state, and if its morph differs, it moves to the + alive state whose morph it matches best instead. ``resolve=False`` always truncates. + """ + history_to_state = {h: index for index in alive for h in states[index]} + edges: dict[int, dict[Any, dict[int, set[History]]]] = {} + for index in alive: + by_symbol: dict[Any, dict[int, set[History]]] = defaultdict(lambda: defaultdict(set)) + for history in states[index]: + for symbol, count in counts.next_counts.get(history, Counter()).items(): + if count == 0: + continue + extended = (*history, symbol) + if len(extended) <= Lmax: + target = history_to_state.get(extended) + else: + target = history_to_state.get(extended[1:]) + if ( + resolve + and counts.next_counts.get(extended) + and ( + target is None or morphs_differ(counts, states[target], {extended}, alpha=alpha, test=test) + ) + ): + best_score = float("inf") + for candidate in sorted(alive): + if morphs_differ(counts, states[candidate], {extended}, alpha=alpha, test=test): + continue + score = morph_test_score(counts, states[candidate], {extended}, test=test) + if score < best_score: + target, best_score = candidate, score + if target is not None: + by_symbol[symbol][target].add(history) + edges[index] = by_symbol + return edges + + +def _recurrent_states(edges: dict[int, dict[Any, dict[int, set[History]]]]) -> list[set[int]]: + """Closed communicating classes (with at least one edge) of the state graph.""" + import networkx as nx + + graph = nx.DiGraph() + graph.add_nodes_from(edges) + for source, by_symbol in edges.items(): + for targets in by_symbol.values(): + graph.add_edges_from((source, target) for target in targets) + condensed = nx.condensation(graph) + classes = [] + for node in condensed: + members = set(condensed.nodes[node]["members"]) + if condensed.out_degree(node) == 0 and graph.subgraph(members).number_of_edges() > 0: + classes.append(members) + return classes + + +def _suffix_determinize( + states: list[set[History]], + counts: SuffixCounts, + alive: set[int], + *, + Lmax: int, + alpha: float, + test: Literal["g", "chi2", "tv"], + resolve: bool = True, +) -> tuple[list[set[History]], set[int]]: + """Split alive states until each (state, symbol) pair has a single alive successor. + + Successors in pruned (transient) states are ignored, as in :cite:`Shalizi2002`. + """ + states = [set(h) for h in states] + alive = set(alive) + while True: + edges = _suffix_edges(states, counts, alive, Lmax=Lmax, alpha=alpha, test=test, resolve=resolve) + split = None + for index in sorted(alive): + for symbol in sorted(edges[index], key=repr): + if len(edges[index][symbol]) > 1: + split = (index, symbol) + break + if split: + break + if split is None: + return states, alive + index, symbol = split + groups = sorted(edges[index][symbol].values(), key=lambda g: (-len(g), sorted(map(repr, g)))) + for group in groups[1:]: + states[index] -= group + states.append(set(group)) + alive.add(len(states) - 1) + + +def _suffix_machine( + states: list[set[History]], + counts: SuffixCounts, + sequence: Sequence[Any], + alive: set[int], + *, + Lmax: int, + alpha: float, + test: Literal["g", "chi2", "tv"], + resolve: bool = True, +) -> EpsilonMachine: + """Build the ε-machine on the most-visited recurrent class of the alive states.""" + edges = _suffix_edges(states, counts, alive, Lmax=Lmax, alpha=alpha, test=test, resolve=resolve) + history_to_state = {h: index for index in alive for h in states[index]} + + visits: Counter[int] = Counter() + seq = tuple(sequence) + for t in range(len(seq) + 1): + for length in range(min(t, Lmax), -1, -1): + state = history_to_state.get(seq[t - length : t]) + if state is not None: + visits[state] += 1 + break + + classes = _recurrent_states(edges) + if not classes: + raise StochasticValidationError("no recurrent inferred states; the sample is too short for this Lmax") + keep = max(classes, key=lambda members: (sum(visits[s] for s in members), -min(members))) + + labels = {state: f"s{rank}" for rank, state in enumerate(sorted(keep))} + transitions = TransitionGraph() + for state in sorted(keep): + transitions.add_state(labels[state]) + for state in sorted(keep): + longest = max(len(h) for h in states[state]) + observed = _observed_counts_for_morph(counts, {h for h in states[state] if len(h) == longest}) + weights = { + symbol: (next(iter(targets)), float(observed.get(symbol, 0))) + for symbol, targets in edges[state].items() + if observed.get(symbol, 0) > 0 + } + if not weights: + weights = { + symbol: (next(iter(targets)), float(sum(len(h) for h in targets.values()))) + for symbol, targets in edges[state].items() + } + total = sum(weight for _, weight in weights.values()) + for symbol, (target, weight) in sorted(weights.items(), key=lambda item: repr(item[0])): + transitions.add_transition( + labels[state], labels[target], **{ATTR_PROB: weight / total, ATTR_EMISSION: symbol} + ) + + kept_visits = {state: visits[state] for state in keep if visits[state] > 0} + total_visits = float(sum(kept_visits.values())) + initial = ( + {labels[state]: count / total_visits for state, count in kept_visits.items()} + if total_visits > 0 + else {labels[min(keep)]: 1.0} + ) + machine = EpsilonMachine( + graph=transitions, + initial_distribution=initial, + observation_alphabet=frozenset(counts.alphabet), + ) + machine.validate() + return machine + + def cssr( sequence: Sequence[Any], *, alphabet: Sequence[Any] | None = None, Lmax: int | None = None, - alpha: float = 0.05, + alpha: float = 0.01, test: Literal["g", "chi2", "tv"] = "g", min_count: int = 5, ) -> EpsilonMachine: - """Reconstruct an ε-machine by Causal-State Splitting Reconstruction (CSSR).""" + """Reconstruct an ε-machine by Causal-State Splitting Reconstruction :cite:`Shalizi2004`. + + Suffixes are grown one symbol into the past up to length ``Lmax`` and grouped by + their next-symbol distributions (homogenization), then states are split until + every transition is deterministic (determinization). The result is restricted + to its most-visited closed class, so it is always a valid recurrent machine. + + Parameters + ---------- + sequence + Observed symbols. + alphabet + Symbol alphabet; defaults to the symbols in ``sequence``. + Lmax + Longest suffix considered; by default a third of ``log_k len(sequence)``, + between 1 and 10. It should be at least the synchronization length of the + source (for a Markov source, its order). Larger values run many more + significance tests, and some of them split states by chance. + alpha + Significance level of each morph-equality test. The worked example of + :cite:`Shalizi2002` uses 0.01; smaller values guard against spurious states + when ``Lmax`` is large. + test + ``"g"`` (G-test), ``"chi2"``, or ``"tv"`` (total-variation threshold). + min_count + Suffixes seen fewer than this many times are not tested or placed in a state. + + Notes + ----- + A process that is not exactly synchronizable (no finite past determines its + state, such as :func:`~sofic.examples.processes.ABC`) has no finite-``Lmax`` + reconstruction. CSSR then returns more states than the ε-machine, with an + entropy rate that approaches the true one from above as ``Lmax`` grows. + """ seq = tuple(sequence) if len(seq) < 2: raise ValueError("sequence must contain at least two symbols") alphabet_size = len(set(seq)) if alphabet is None else len(tuple(alphabet)) - max_length = Lmax if Lmax is not None else _default_lmax(len(seq), alphabet_size, min_count) + max_length = Lmax if Lmax is not None else _cssr_default_lmax(len(seq), alphabet_size) + if max_length < 0: + raise ValueError("Lmax must be non-negative") counts = SuffixCounts.from_sequence(seq, alphabet=alphabet, max_length=max_length + 1) - states, history_to_state = _cssr_homogenize( - counts, - Lmax=max_length, - alpha=alpha, - test=test, - ) - states = _cssr_determinize(states, history_to_state, counts, length=max_length) - states = _merge_similar_states(states, history_to_state, counts, alpha=alpha, test=test) - # Re-determinize: morph-only merging can fuse states with incompatible - # successors, so restore unifilarity before building the machine. - states = _cssr_determinize(states, history_to_state, counts, length=max_length) - states = _drop_transient_states(states, history_to_state, counts, length=max_length) - history_to_state = {history: state_id for state_id, histories in states.items() for history in histories} - return _counts_to_mealy(states, counts, history_to_state, seq, length=max_length) + homogeneous = _suffix_homogenize(counts, Lmax=max_length, alpha=alpha, test=test, min_count=min_count) + + def reconstruct(resolve: bool) -> EpsilonMachine: + everything = set(range(len(homogeneous))) + edges = _suffix_edges(homogeneous, counts, everything, Lmax=max_length, alpha=alpha, test=test, resolve=resolve) + alive = set().union(*_recurrent_states(edges)) or everything + states, alive = _suffix_determinize( + homogeneous, counts, alive, Lmax=max_length, alpha=alpha, test=test, resolve=resolve + ) + return _suffix_machine(states, counts, seq, alive, Lmax=max_length, alpha=alpha, test=test, resolve=resolve) + + machine = reconstruct(resolve=True) + # Resolving truncated successors needs Lmax at least the synchronization length. When it + # is shorter, resolution can close off a state that never emits some observed symbol. + if {t.data[ATTR_EMISSION] for t in machine.transitions()} < set(seq): + machine = reconstruct(resolve=False) + return machine def _morph_distance( diff --git a/sofic/generators/hmm_inference.py b/sofic/generators/hmm_inference.py index 7be8e8b..eb86703 100644 --- a/sofic/generators/hmm_inference.py +++ b/sofic/generators/hmm_inference.py @@ -26,7 +26,12 @@ def _as_mealy_hmm(hmm: HiddenMarkovModel) -> Any: def _emission_transition_tensors_from_mealy( hmm: Any, ) -> tuple[np.ndarray, dict[Any, np.ndarray]]: - """Return initial vector ``pi`` and symbol -> joint transition matrices.""" + """Return initial vector ``pi`` and symbol -> joint transition matrices. + + Symbols are keyed in a fixed (``repr``-sorted) order, so seeded sampling is + reproducible across interpreter runs. A model without an initial + distribution starts from its stationary distribution. + """ from sofic.generators.prob import as_prob, has_symbolic, zeros idx = hmm.reindex() @@ -36,10 +41,14 @@ def _emission_transition_tensors_from_mealy( symbolic = has_symbolic(edge_probs) or has_symbolic(init_probs) pi = zeros((n,), symbolic=symbolic) - for state, mass in hmm.initial_distribution.items(): - pi[idx.index(state)] = as_prob(mass) - - symbols: set[Any] = set(hmm.observation_alphabet) + if hmm.initial_distribution: + for state, mass in hmm.initial_distribution.items(): + pi[idx.index(state)] = as_prob(mass) + elif n: + pi = np.asarray(hmm.stationary_distribution(), dtype=object if symbolic else float) + + emissions = {transition.data.get(ATTR_EMISSION) for transition in hmm.transitions()} - {None} + symbols = sorted(set(hmm.observation_alphabet) | emissions, key=repr) joint: dict[Any, np.ndarray] = {symbol: zeros((n, n), symbolic=symbolic) for symbol in symbols} for transition in hmm.transitions(): diff --git a/sofic/generators/mealy.py b/sofic/generators/mealy.py index 24e300d..82b7ef3 100644 --- a/sofic/generators/mealy.py +++ b/sofic/generators/mealy.py @@ -35,10 +35,13 @@ def add_transition(self, source: Hashable, target: Hashable, symbol: Any, prob: """Add an edge carrying joint emission probability ``P(target, symbol | source)``. ``prob`` may be a Python float or an exact sympy expression (see - :mod:`sofic.generators.prob`). + :mod:`sofic.generators.prob`). ``symbol`` is added to the observation + alphabet if it is not already there. """ from sofic.generators.prob import as_prob + if symbol not in self.observation_alphabet: + self.observation_alphabet = self.observation_alphabet | {symbol} return self.graph.add_transition( source, target, diff --git a/sofic/generators/prob.py b/sofic/generators/prob.py index 8e8ed8e..d3dbe72 100644 --- a/sofic/generators/prob.py +++ b/sofic/generators/prob.py @@ -2,6 +2,7 @@ from __future__ import annotations +import sys from collections.abc import Iterable, Sequence from typing import Any @@ -24,9 +25,10 @@ def is_symbolic(value: Any) -> bool: Exact rationals / integers and expressions with free symbols are symbolic. ``sympy.Float`` is treated as numeric and will be coerced to ``float``. """ - try: - import sympy - except ImportError: + # A sympy value can only exist once sympy is imported. Importing here instead + # would retry a failed import (a filesystem search) on every call without sympy. + sympy = sys.modules.get("sympy") + if sympy is None: return False if not isinstance(value, sympy.Expr): return False @@ -50,13 +52,6 @@ def as_prob(value: Any) -> Prob: return float(value) if isinstance(value, int): return value - try: - import sympy - - if isinstance(value, sympy.Basic): - return float(value) - except ImportError: - pass return float(value) diff --git a/tests/test_epsilon_inference.py b/tests/test_epsilon_inference.py index b025d24..be31dc6 100644 --- a/tests/test_epsilon_inference.py +++ b/tests/test_epsilon_inference.py @@ -204,3 +204,58 @@ def test_from_sequence_spectral_dispatch(rng: np.random.Generator): def test_from_sequence_unknown_method(): with pytest.raises(ValueError, match="unknown inference method"): EpsilonMachine.from_sequence([0, 1, 0], method="nsd") + + +@pytest.mark.parametrize( + ("name", "Lmax"), + [("Even", 3), ("Even", 5), ("GoldenMean", 3), ("Nemo", 4), ("RkGM", 5)], +) +def test_cssr_recovers_synchronizable_processes(name: str, Lmax: int): + """Regression: appended (not prepended) suffixes and untruncated successors + dropped edges, so these raised StochasticValidationError or returned h_mu = 0.""" + from sofic.examples import processes + + oracle = processes.RkGM(5, 3) if name == "RkGM" else getattr(processes, name)() + observations, _ = sample(oracle, 20000, np.random.default_rng(5)) + inferred = cssr(observations, Lmax=Lmax, alpha=0.001) + inferred.validate() + assert inferred.is_unifilar() + assert len(list(inferred.states())) == len(list(oracle.states())) + assert inferred.entropy_rate() == pytest.approx(oracle.entropy_rate(), abs=0.02) + + +def test_cssr_even_process_ignores_truncated_nonsynchronizing_suffix(): + """At Lmax = 3 the successor of ``011`` on ``1`` truncates to the ambiguous ``111``.""" + oracle = even_process(0.5) + observations, _ = sample(oracle, 20000, np.random.default_rng(5)) + inferred = cssr(observations, Lmax=3, alpha=0.01) + assert _signatures_isomorphic(inferred, oracle, prob_tol=0.03) + + +def test_cssr_default_lmax_does_not_oversplit(): + for oracle, n_states in [(even_process(0.5), 2), (bernoulli(0.3), 1)]: + observations, _ = sample(oracle, 20000, np.random.default_rng(7)) + inferred = cssr(observations) + inferred.validate() + assert len(list(inferred.states())) == n_states + + +def test_cssr_short_lmax_still_emits_every_symbol(): + """Lmax below the Markov order cannot recover RkGM(5, 3), but must not collapse to a trap state.""" + from sofic.examples import processes + + observations, _ = sample(processes.RkGM(5, 3), 20000, np.random.default_rng(5)) + inferred = cssr(observations, Lmax=3, alpha=0.001) + inferred.validate() + assert {t.data[ATTR_EMISSION] for t in inferred.transitions()} == {"0", "1"} + assert inferred.entropy_rate() > 0.0 + + +def test_cssr_non_synchronizable_process_returns_valid_machine(): + from sofic.examples import processes + + oracle = processes.ABC() + observations, _ = sample(oracle, 20000, np.random.default_rng(5)) + inferred = cssr(observations, Lmax=4, alpha=0.001) + inferred.validate() + assert inferred.entropy_rate() >= oracle.entropy_rate() - 0.02 diff --git a/tests/test_hmm_inference.py b/tests/test_hmm_inference.py index 9e1c643..2cf6756 100644 --- a/tests/test_hmm_inference.py +++ b/tests/test_hmm_inference.py @@ -322,3 +322,28 @@ def test_hmm_methods_delegate_to_inference_functions(): assert np.allclose(gm.observed_information(obs), observed_information(gm, obs)) _fitted, trace = gm.baum_welch(obs) assert trace + + +def test_seeded_sample_is_reproducible_across_hash_seeds(): + """Regression: symbol order came from iterating a frozenset, so a seeded sample + changed with PYTHONHASHSEED.""" + import os + import subprocess + import sys + + script = ( + "import numpy as np\n" + "from sofic.examples.processes import Nemo\n" + "print(''.join(Nemo().sample(200, rng=np.random.default_rng(5))[0]))\n" + ) + outputs = { + subprocess.run( + [sys.executable, "-W", "ignore", "-c", script], + capture_output=True, + text=True, + check=True, + env={**os.environ, "PYTHONHASHSEED": seed}, + ).stdout + for seed in ("1", "2", "3") + } + assert len(outputs) == 1 diff --git a/tests/test_mealy_hmm.py b/tests/test_mealy_hmm.py index 5e41341..d95ea62 100644 --- a/tests/test_mealy_hmm.py +++ b/tests/test_mealy_hmm.py @@ -83,3 +83,19 @@ def test_entropy_rate_requires_unifilar_presentation(): with pytest.raises(NotImplementedError, match="unifilar"): hmm.entropy_rate() + + +def test_mealy_without_alphabet_or_initial_distribution(): + """Regression: symbols missing from the alphabet raised KeyError, and an empty + initial distribution made sampling fail on NaN and every word probability 0.""" + import numpy as np + + hmm = MealyHMM() + for source, target, prob in [("a", "a", 0.9), ("a", "b", 0.1), ("b", "a", 0.2), ("b", "b", 0.8)]: + hmm.add_transition(source, target, target, prob) + assert hmm.observation_alphabet == frozenset({"a", "b"}) + hmm.validate() + symbols, _ = hmm.sample(50, rng=np.random.default_rng(0)) + assert len(symbols) == 50 + assert hmm.word_probability("ab") == pytest.approx(2 / 3 * 0.1) + assert hmm.entropy_rate() == pytest.approx(0.5533064, abs=1e-6) diff --git a/tests/test_prob.py b/tests/test_prob.py new file mode 100644 index 0000000..e84d210 --- /dev/null +++ b/tests/test_prob.py @@ -0,0 +1,25 @@ +"""Tests for probability scalar helpers.""" + +from __future__ import annotations + +import builtins + +from sofic.generators import prob + + +def test_numeric_checks_never_import_sympy(monkeypatch): + """Regression: without sympy installed, every check retried the failed import, + which made entropy_rate() on a 127-state machine take seconds.""" + calls = [] + real_import = builtins.__import__ + + def spy(name, *args, **kwargs): + if name == "sympy": + calls.append(name) + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", spy) + assert not prob.is_symbolic(0.5) + assert not prob.has_symbolic([0.1, 0.9]) + assert prob.as_prob(0.25) == 0.25 + assert calls == [] From d563094ef634f01552dcba7d0ab02a69e2436f82 Mon Sep 17 00:00:00 2001 From: Ryan James Date: Sun, 27 Sep 2026 20:46:43 -0600 Subject: [PATCH 2/2] Fix successor handling in subtree merging, stack CSSR, and transCSSR - subtree_merge: use the CSSR successor rule (extend, truncate at L with resolution) and keep the most-visited recurrent class; delta = 0 now means a G-test at 0.01 instead of a fixed 1e-3 tolerance, which made the default raise StochasticValidationError on sampled data. - stack_cssr: seed a root per observed stack and grow suffixes into the past (forward growth only reached stacks of depth <= Lmax); compare morphs with return symbols pooled; match return edges only to observed calls; weight each symbol by its frequency where legal; transient removal now uses the stack successor (it was a no-op); consistent max_stack_depth; default Lmax follows cssr. - HiddenMarkovStackModel.word_probability: renormalize each step instead of pruning below an absolute 1e-15, which zeroed every word longer than ~50. - transcssr: grow joint suffixes into the past, truncate and resolve successors, drop transient states before determinizing (Delay(2) gave 5-19 states and forbade valid pairs). - Inline the G statistic (identical to scipy's chi2_contingency, including Yates' correction) and cache critical values; stack subtree merging on the Motzkin shift went from 11 s to 1.6 s. Co-authored-by: Cursor --- docs/generators/epsilon_inference.rst | 5 +- docs/generators/stack_inference.rst | 7 + sofic/generators/epsilon_inference.py | 200 +++------ .../epsilon_transducer_inference.py | 403 +++++++++--------- sofic/generators/stack_hmm.py | 11 +- sofic/generators/stack_inference.py | 229 ++++++---- tests/test_epsilon_inference.py | 14 + tests/test_epsilon_transducer_inference.py | 30 ++ tests/test_stack_hmm.py | 13 + tests/test_stack_inference.py | 25 ++ 10 files changed, 510 insertions(+), 427 deletions(-) diff --git a/docs/generators/epsilon_inference.rst b/docs/generators/epsilon_inference.rst index 6a12d26..b263141 100644 --- a/docs/generators/epsilon_inference.rst +++ b/docs/generators/epsilon_inference.rst @@ -66,8 +66,9 @@ Subtree merging Subtree merging :cite:`CrutchfieldYoung1989` clusters histories with statistically equivalent next-symbol distributions (metric tolerance ``delta``), then determinizes -to a unifilar presentation. With ``delta=0``, morphs are compared up to a small -numerical tolerance for finite-sample estimates. +to a unifilar presentation. With ``delta=0``, two morphs are equivalent unless a +G-test at significance 0.01 tells them apart, a tolerance that scales with the +sample. Transitions follow the same successor rule as CSSR. .. autofunction:: subtree_merge diff --git a/docs/generators/stack_inference.rst b/docs/generators/stack_inference.rst index 05a286b..bc9395d 100644 --- a/docs/generators/stack_inference.rst +++ b/docs/generators/stack_inference.rst @@ -41,6 +41,13 @@ Two families are provided: Histories are counted with a bounded stack depth, so ``max_stack_depth`` caps the configurations considered during reconstruction. +Stack CSSR runs flat CSSR over ``(suffix, stack)`` configurations. Every observed +stack seeds its own root, and suffixes grow into the past with the stack fixed. +When morphs are compared, all return symbols count as one event: which return +can follow is decided by the stack top through matched call-return pairs, not by +the finite control. Return edges are matched only to calls observed to close +them. + API === diff --git a/sofic/generators/epsilon_inference.py b/sofic/generators/epsilon_inference.py index 7a7caf0..37adf36 100644 --- a/sofic/generators/epsilon_inference.py +++ b/sofic/generators/epsilon_inference.py @@ -11,6 +11,7 @@ from collections import Counter, defaultdict from collections.abc import Callable, Iterable, Sequence from dataclasses import dataclass, field +from functools import lru_cache from typing import Any, ClassVar, Literal import numpy as np @@ -150,6 +151,30 @@ def _contingency_rows( return table +def _g_statistic(table: np.ndarray) -> float | None: + """G-test statistic of a contingency table, as ``scipy.stats.chi2_contingency`` computes it. + + Includes Yates' correction for one degree of freedom, and returns ``None`` if an + expected count is zero. Inlined because the test runs once per pair of + histories, and the general scipy routine dominated inference time. + """ + expected = table.sum(axis=1, keepdims=True) * table.sum(axis=0, keepdims=True) / table.sum() + if np.any(expected == 0): + return None + observed = table + if (table.shape[0] - 1) * (table.shape[1] - 1) == 1: + diff = expected - observed + observed = observed + np.sign(diff) * np.minimum(0.5, np.abs(diff)) + with np.errstate(divide="ignore", invalid="ignore"): + terms = np.where(observed > 0, observed * np.log(observed / expected), 0.0) + return 2.0 * float(terms.sum()) + + +@lru_cache(maxsize=256) +def _chi2_critical(alpha: float, dof: int) -> float: + return float(stats.chi2.ppf(1.0 - alpha, dof)) + + def morphs_differ( counts: SuffixCounts, left_histories: set[History], @@ -170,16 +195,10 @@ def morphs_differ( if table is None: return False if test == "g": - try: - with np.errstate(invalid="ignore", divide="ignore"): - statistic, _p_value, _dof, expected = stats.chi2_contingency(table, lambda_="log-likelihood") - except ValueError: + statistic = _g_statistic(table) + if statistic is None or not np.isfinite(statistic): return False - if not np.isfinite(statistic) or np.any(expected == 0): - return False - dof = max(1, table.shape[1] - 1) - critical = float(stats.chi2.ppf(1.0 - alpha, dof)) - return float(statistic) > critical + return statistic > _chi2_critical(alpha, max(1, table.shape[1] - 1)) try: statistic, p_value, _dof, expected = stats.chi2_contingency(table) except ValueError: @@ -206,9 +225,8 @@ def morph_test_score( return 0.0 try: if test == "g": - with np.errstate(invalid="ignore", divide="ignore"): - statistic, _p, _dof, _expected = stats.chi2_contingency(table, lambda_="log-likelihood") - return float(statistic) if np.isfinite(statistic) else 0.0 + statistic = _g_statistic(table) + return statistic if statistic is not None and np.isfinite(statistic) else 0.0 statistic, _p, _dof, _expected = stats.chi2_contingency(table) return float(statistic) except ValueError: @@ -386,6 +404,7 @@ def _drop_transient_states( counts: SuffixCounts, *, length: int, + successor_fn: Callable[[History, Any], History] = _grow_history, ) -> dict[int, set[History]]: """Keep only states in bottom strongly connected components.""" import networkx as nx @@ -396,8 +415,7 @@ def _drop_transient_states( for symbol in counts.alphabet: if counts.next_counts.get(history, Counter()).get(symbol, 0) == 0: continue - child = history + (symbol,) - target = history_to_state.get(child) + target = history_to_state.get(successor_fn(history, symbol)) if target is None: continue successors[state_id][symbol].add(target) @@ -432,84 +450,6 @@ def _drop_transient_states( return {state_id: histories for state_id, histories in states.items() if state_id in recurrent} -def _empirical_state_visits( - sequence: Sequence[Any], - history_to_state: dict[History, int], - *, - length: int, -) -> Counter[int]: - """Count how often each causal state is occupied along ``sequence``. - - Every time step belongs to exactly one causal state — the one keyed by the - *longest* available suffix (up to ``length``). Counting each nested suffix - (as an earlier version did) over-weights short-history states and skews the - reconstructed ``initial_distribution`` away from the occupation/stationary law. - """ - visits: Counter[int] = Counter() - seq = tuple(sequence) - for t in range(len(seq)): - for hist_len in range(min(t, length), -1, -1): - history = seq[t - hist_len : t] - state = history_to_state.get(history) - if state is not None: - visits[state] += 1 - break - return visits - - -def _counts_to_mealy( - states: dict[int, set[History]], - counts: SuffixCounts, - history_to_state: dict[History, int], - sequence: Sequence[Any], - *, - length: int, -) -> EpsilonMachine: - visits = _empirical_state_visits(sequence, history_to_state, length=length) - if not visits: - raise StochasticValidationError("no empirical state visits") - - graph = TransitionGraph() - state_labels = {state_id: f"s{state_id}" for state_id in states} - for label in state_labels.values(): - graph.add_state(label) - - for state_id, histories in states.items(): - label = state_labels[state_id] - morph = counts.state_morph(histories) - for symbol in counts.alphabet: - prob = morph[symbol] - if prob <= 0.0: - continue - emitting = [ - history for history in histories if counts.next_counts.get(history, Counter()).get(symbol, 0) > 0 - ] - if not emitting: - continue - child_histories = {history + (symbol,) for history in emitting} - targets = {history_to_state.get(child) for child in child_histories} - targets.discard(None) - if not targets: - continue - if len(targets) > 1: - raise StochasticValidationError(f"non-unifilar inferred transition from {label!r} on {symbol!r}") - target_label = state_labels[next(iter(targets))] - graph.add_transition(label, target_label, **{ATTR_PROB: prob, ATTR_EMISSION: symbol}) - - total_visits = float(sum(visits.values())) - initial = {state_labels[state_id]: visits[state_id] / total_visits for state_id in states if visits[state_id] > 0} - if not initial: - initial = {state_labels[next(iter(states))]: 1.0} - - machine = EpsilonMachine( - graph=graph, - initial_distribution=initial, - observation_alphabet=frozenset(counts.alphabet), - ) - machine.validate() - return machine - - def _cssr_default_lmax(n: int, alphabet_size: int) -> int: """A third of ``log_k n``, between 1 and 10. @@ -744,6 +684,32 @@ def _suffix_machine( return machine +def _suffix_reconstruct( + states: list[set[History]], + counts: SuffixCounts, + sequence: Sequence[Any], + *, + Lmax: int, + alpha: float, + test: Literal["g", "chi2", "tv"], +) -> EpsilonMachine: + """Prune transient states, determinize, and build the machine from homogeneous ``states``.""" + + def reconstruct(resolve: bool) -> EpsilonMachine: + everything = set(range(len(states))) + edges = _suffix_edges(states, counts, everything, Lmax=Lmax, alpha=alpha, test=test, resolve=resolve) + alive = set().union(*_recurrent_states(edges)) or everything + split, alive = _suffix_determinize(states, counts, alive, Lmax=Lmax, alpha=alpha, test=test, resolve=resolve) + return _suffix_machine(split, counts, sequence, alive, Lmax=Lmax, alpha=alpha, test=test, resolve=resolve) + + machine = reconstruct(resolve=True) + # Resolving truncated successors needs Lmax at least the synchronization length. When it + # is shorter, resolution can close off a state that never emits some observed symbol. + if {t.data[ATTR_EMISSION] for t in machine.transitions()} < set(sequence): + machine = reconstruct(resolve=False) + return machine + + def cssr( sequence: Sequence[Any], *, @@ -797,22 +763,11 @@ def cssr( counts = SuffixCounts.from_sequence(seq, alphabet=alphabet, max_length=max_length + 1) 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) - def reconstruct(resolve: bool) -> EpsilonMachine: - everything = set(range(len(homogeneous))) - edges = _suffix_edges(homogeneous, counts, everything, Lmax=max_length, alpha=alpha, test=test, resolve=resolve) - alive = set().union(*_recurrent_states(edges)) or everything - states, alive = _suffix_determinize( - homogeneous, counts, alive, Lmax=max_length, alpha=alpha, test=test, resolve=resolve - ) - return _suffix_machine(states, counts, seq, alive, Lmax=max_length, alpha=alpha, test=test, resolve=resolve) - machine = reconstruct(resolve=True) - # Resolving truncated successors needs Lmax at least the synchronization length. When it - # is shorter, resolution can close off a state that never emits some observed symbol. - if {t.data[ATTR_EMISSION] for t in machine.transitions()} < set(seq): - machine = reconstruct(resolve=False) - return machine +#: Significance level used by subtree merging when ``delta = 0`` and to resolve truncated successors. +_SUBTREE_ALPHA = 0.01 def _morph_distance( @@ -834,11 +789,9 @@ def _morphs_equivalent( *, delta: float, ) -> bool: - left_morph = counts.morph(left) - right_morph = counts.morph(right) if delta > 0.0: return _morph_distance(counts, left, right, delta=delta) <= delta - return all(np.isclose(left_morph[symbol], right_morph[symbol], rtol=0.0, atol=1e-3) for symbol in counts.alphabet) + return not morphs_differ(counts, {left}, {right}, alpha=_SUBTREE_ALPHA, test="g") def _cluster_histories_by_morph( @@ -885,7 +838,13 @@ def subtree_merge( delta: float = 0.0, alphabet: Sequence[Any] | None = None, ) -> EpsilonMachine: - """Reconstruct an ε-machine by merging depth-``L`` subtrees (Crutchfield--Young).""" + """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 G-test at + significance 0.01 tells them apart. The clusters are then determinized as in + :func:`cssr`. + """ if L < 0: raise ValueError("L must be non-negative") seq = tuple(sequence) @@ -895,26 +854,9 @@ def subtree_merge( histories = {history for history in counts.history_counts if len(history) <= L} histories.add(()) + states = list(_cluster_histories_by_morph(counts, histories, delta=delta).values()) - states = _cluster_histories_by_morph(counts, histories, delta=delta) - history_to_state = {history: state_id for state_id, members in states.items() for history in members} - - states = _cssr_determinize(states, history_to_state, counts, length=L) - history_to_state = { - history: state_id for state_id, histories_in_state in states.items() for history in histories_in_state - } - states = _merge_similar_states(states, history_to_state, counts, alpha=0.05, test="tv") - history_to_state = { - history: state_id for state_id, histories_in_state in states.items() for history in histories_in_state - } - # Re-determinize: morph-only merging can fuse states with incompatible - # successors, so restore unifilarity before building the machine. - states = _cssr_determinize(states, history_to_state, counts, length=L) - states = _drop_transient_states(states, history_to_state, counts, length=L) - history_to_state = { - history: state_id for state_id, histories_in_state in states.items() for history in histories_in_state - } - return _counts_to_mealy(states, counts, history_to_state, seq, length=L) + return _suffix_reconstruct(states, counts, seq, Lmax=L, alpha=_SUBTREE_ALPHA, test="g") def spectral( diff --git a/sofic/generators/epsilon_transducer_inference.py b/sofic/generators/epsilon_transducer_inference.py index 9f9a6b5..b782b8d 100644 --- a/sofic/generators/epsilon_transducer_inference.py +++ b/sofic/generators/epsilon_transducer_inference.py @@ -13,7 +13,7 @@ from __future__ import annotations from collections import Counter, defaultdict -from collections.abc import Sequence +from collections.abc import Iterable, Sequence from dataclasses import dataclass, field from typing import Any, Literal @@ -185,6 +185,17 @@ def _table_significant(table: np.ndarray, *, alpha: float, test: Literal["g", "c return float(p_value) < alpha +def _state_aggregate(counts: JointSuffixCounts, histories: Iterable[JointHistory]) -> StateAggregate: + aggregate: StateAggregate = {} + for history in histories: + _merge_aggregate(aggregate, _history_aggregate(counts, history)) + return aggregate + + +def _observed(counts: JointSuffixCounts, history: JointHistory) -> int: + return sum(sum(counter.values()) for counter in counts.next_counts.get(history, {}).values()) + + def _homogenize( counts: JointSuffixCounts, *, @@ -192,122 +203,50 @@ def _homogenize( alpha: float, test: Literal["g", "chi2"], min_count: int, -) -> tuple[dict[int, set[JointHistory]], dict[JointHistory, int]]: - in_alpha = counts.input_alphabet - out_alpha = counts.output_alphabet - states: dict[int, set[JointHistory]] = {0: {()}} - state_agg: dict[int, StateAggregate] = {0: _history_aggregate(counts, ())} - history_to_state: dict[JointHistory, int] = {(): 0} - next_state_id = 1 - - for _length in range(Lmax + 1): - for state_id in sorted(states): - for history in list(states[state_id]): - for pair in _observed_pairs(counts, history): - child = (*history, pair) - if child in history_to_state or counts.history_counts.get(child, 0) == 0: +) -> list[set[JointHistory]]: + """transCSSR homogenization: grow joint suffixes one ``(input, output)`` pair into the past. + + As in flat CSSR, a child ``p h`` of suffix ``h`` stays in its parent's state + unless its conditional output law differs; histories seen fewer than + ``min_count`` times stay with their parent. + """ + in_alpha, out_alpha = counts.input_alphabet, counts.output_alphabet + states: list[set[JointHistory]] = [{()}] + aggregates: list[StateAggregate] = [_history_aggregate(counts, ())] + pairs = _pairs_from(counts) + + def differ(left: StateAggregate, right: StateAggregate) -> bool: + return aggregates_differ( + left, right, input_alphabet=in_alpha, output_alphabet=out_alpha, alpha=alpha, test=test + ) + + for length in range(Lmax): + for parent_id in range(len(states)): + for history in sorted((h for h in states[parent_id] if len(h) == length), key=repr): + for pair in pairs: + child = (pair, *history) + if _observed(counts, child) == 0: continue child_agg = _history_aggregate(counts, child) - if counts.history_counts.get(child, 0) < min_count: - # Too rare to split reliably; inherit the parent's causal state. - states[state_id].add(child) - history_to_state[child] = state_id - _merge_aggregate(state_agg[state_id], child_agg) - continue - if aggregates_differ( - state_agg[state_id], - child_agg, - input_alphabet=in_alpha, - output_alphabet=out_alpha, - alpha=alpha, - test=test, - ): - best_state: int | None = None - best_score = float("inf") - for candidate_id, candidate_agg in state_agg.items(): - if aggregates_differ( - candidate_agg, - child_agg, - input_alphabet=in_alpha, - output_alphabet=out_alpha, - alpha=alpha, - test=test, - ): + target = parent_id + if _observed(counts, child) >= min_count and differ(aggregates[parent_id], child_agg): + best_id, best_score = None, float("inf") + for candidate_id, candidate_agg in enumerate(aggregates): + if candidate_id == parent_id or differ(candidate_agg, child_agg): continue score = _aggregate_score( - candidate_agg, - child_agg, - input_alphabet=in_alpha, - output_alphabet=out_alpha, + candidate_agg, child_agg, input_alphabet=in_alpha, output_alphabet=out_alpha ) if score < best_score: - best_score = score - best_state = candidate_id - if best_state is None: - best_state = next_state_id - states[next_state_id] = set() - state_agg[next_state_id] = {} - next_state_id += 1 - states[best_state].add(child) - history_to_state[child] = best_state - _merge_aggregate(state_agg[best_state], child_agg) - else: - states[state_id].add(child) - history_to_state[child] = state_id - _merge_aggregate(state_agg[state_id], child_agg) - return states, history_to_state - - -def _observed_pairs(counts: JointSuffixCounts, history: JointHistory) -> list[tuple[Any, Any]]: - by_input = counts.next_counts.get(history) - if by_input is None: - return [] - pairs: list[tuple[Any, Any]] = [] - for input_symbol, counter in by_input.items(): - for output_symbol in counter: - pairs.append((input_symbol, output_symbol)) - return pairs - - -def _determinize( - states: dict[int, set[JointHistory]], - history_to_state: dict[JointHistory, int], - counts: JointSuffixCounts, -) -> dict[int, set[JointHistory]]: - current = {state_id: set(histories) for state_id, histories in states.items()} - next_state_id = (max(current) + 1) if current else 0 - changed = True - while changed: - changed = False - for state_id in sorted(current): - histories = current[state_id] - if len(histories) <= 1: - continue - for pair in _pairs_from(counts): - buckets: dict[int, set[JointHistory]] = defaultdict(set) - for history in histories: - if not _history_emits(counts, history, pair): - continue - child = (*history, pair) - target = history_to_state.get(child) - if target is None: - continue - buckets[target].add(history) - if len(buckets) <= 1: - continue - ordered = sorted(buckets.items(), key=lambda item: (-len(item[1]), repr(min(item[1], key=repr)))) - _keep_target, keep_histories = ordered[0] - current[state_id] = keep_histories - for _target, split_histories in ordered[1:]: - current[next_state_id] = split_histories - for history in split_histories: - history_to_state[history] = next_state_id - next_state_id += 1 - changed = True - break - if changed: - break - return current + best_id, best_score = candidate_id, score + if best_id is None: + states.append(set()) + aggregates.append({}) + best_id = len(states) - 1 + target = best_id + states[target].add(child) + _merge_aggregate(aggregates[target], child_agg) + return states def _pairs_from(counts: JointSuffixCounts) -> list[tuple[Any, Any]]: @@ -322,118 +261,182 @@ def _history_emits(counts: JointSuffixCounts, history: JointHistory, pair: tuple return bool(counter) and counter.get(pair[1], 0) > 0 -def _drop_transient( - states: dict[int, set[JointHistory]], - history_to_state: dict[JointHistory, int], +def _edges( + states: list[set[JointHistory]], counts: JointSuffixCounts, -) -> dict[int, set[JointHistory]]: - import networkx as nx + alive: set[int], + *, + Lmax: int, + alpha: float, + test: Literal["g", "chi2"], +) -> dict[int, dict[tuple[Any, Any], dict[int, set[JointHistory]]]]: + """Successor states by ``(input, output)`` pair, with the histories that lead there. - graph = nx.DiGraph() - graph.add_nodes_from(states) - for state_id, histories in states.items(): - for history in histories: + 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`). + """ + in_alpha, out_alpha = counts.input_alphabet, counts.output_alphabet + history_to_state = {h: index for index in alive for h in states[index]} + aggregates = {index: _state_aggregate(counts, states[index]) for index in alive} + edges: dict[int, dict[tuple[Any, Any], dict[int, set[JointHistory]]]] = {} + for index in alive: + by_pair: dict[tuple[Any, Any], dict[int, set[JointHistory]]] = defaultdict(lambda: defaultdict(set)) + for history in states[index]: for pair in _pairs_from(counts): if not _history_emits(counts, history, pair): continue - target = history_to_state.get((*history, pair)) + extended = (*history, pair) + if len(extended) <= Lmax: + target = history_to_state.get(extended) + else: + target = history_to_state.get(extended[1:]) + extended_agg = _history_aggregate(counts, extended) + if extended_agg and ( + target is None + or aggregates_differ( + aggregates[target], + extended_agg, + input_alphabet=in_alpha, + output_alphabet=out_alpha, + alpha=alpha, + test=test, + ) + ): + best_score = float("inf") + for candidate in sorted(alive): + if aggregates_differ( + aggregates[candidate], + extended_agg, + input_alphabet=in_alpha, + output_alphabet=out_alpha, + alpha=alpha, + test=test, + ): + continue + score = _aggregate_score( + aggregates[candidate], extended_agg, input_alphabet=in_alpha, output_alphabet=out_alpha + ) + if score < best_score: + target, best_score = candidate, score if target is not None: - graph.add_edge(state_id, target) - if graph.number_of_edges() == 0: - return states - - recurrent: set[int] = set() - for component in nx.strongly_connected_components(graph): - subgraph = graph.subgraph(component) - has_cycle = subgraph.number_of_edges() > 0 and ( - len(component) > 1 or any(subgraph.has_edge(node, node) for node in component) + by_pair[pair][target].add(history) + edges[index] = by_pair + return edges + + +def _closed_classes(edges: dict[int, dict[tuple[Any, Any], dict[int, set[JointHistory]]]]) -> list[set[int]]: + import networkx as nx + + graph = nx.DiGraph() + graph.add_nodes_from(edges) + for source, by_pair in edges.items(): + for targets in by_pair.values(): + graph.add_edges_from((source, target) for target in targets) + condensed = nx.condensation(graph) + return [ + set(condensed.nodes[node]["members"]) + for node in condensed + if condensed.out_degree(node) == 0 and graph.subgraph(condensed.nodes[node]["members"]).number_of_edges() > 0 + ] + + +def _determinize( + states: list[set[JointHistory]], + counts: JointSuffixCounts, + alive: set[int], + *, + Lmax: int, + alpha: float, + test: Literal["g", "chi2"], +) -> tuple[list[set[JointHistory]], set[int]]: + """Split alive states until each ``(input, output)`` pair has one alive successor.""" + states = [set(h) for h in states] + alive = set(alive) + while True: + edges = _edges(states, counts, alive, Lmax=Lmax, alpha=alpha, test=test) + split = next( + ( + (index, pair) + for index in sorted(alive) + for pair in sorted(edges[index], key=repr) + if len(edges[index][pair]) > 1 + ), + None, ) - if not has_cycle: - continue - if not any(graph.has_edge(v, w) for v in component for w in graph.nodes if w not in component): - recurrent.update(component) - if not recurrent: - return states - return {state_id: histories for state_id, histories in states.items() if state_id in recurrent} + if split is None: + return states, alive + index, pair = split + groups = sorted(edges[index][pair].values(), key=lambda g: (-len(g), sorted(map(repr, g)))) + for group in groups[1:]: + states[index] -= group + states.append(set(group)) + alive.add(len(states) - 1) -def _state_visits( +def _build_transducer( + states: list[set[JointHistory]], + counts: JointSuffixCounts, + alive: set[int], inputs: Sequence[Any], outputs: Sequence[Any], - history_to_state: dict[JointHistory, int], *, - length: int, -) -> Counter[int]: + Lmax: int, + alpha: float, + test: Literal["g", "chi2"], +) -> EpsilonTransducer: + edges = _edges(states, counts, alive, Lmax=Lmax, alpha=alpha, test=test) + history_to_state = {h: index for index in alive for h in states[index]} + visits: Counter[int] = Counter() pairs = tuple(zip(inputs, outputs, strict=True)) - for t in range(len(pairs)): - for hist_len in range(min(t, length), -1, -1): - history = pairs[t - hist_len : t] - state = history_to_state.get(history) + for t in range(len(pairs) + 1): + for hist_len in range(min(t, Lmax), -1, -1): + state = history_to_state.get(pairs[t - hist_len : t]) if state is not None: visits[state] += 1 break - return visits - -def _build_transducer( - states: dict[int, set[JointHistory]], - counts: JointSuffixCounts, - history_to_state: dict[JointHistory, int], - inputs: Sequence[Any], - outputs: Sequence[Any], - *, - length: int, -) -> EpsilonTransducer: - visits = _state_visits(inputs, outputs, history_to_state, length=length) - if not visits: - raise StochasticValidationError("no empirical causal-state visits") + classes = _closed_classes(edges) + if not classes: + raise StochasticValidationError("no recurrent inferred states; the sample is too short for this Lmax") + keep = max(classes, key=lambda members: (sum(visits[s] for s in members), -min(members))) graph = TransitionGraph() - labels = {state_id: f"s{state_id}" for state_id in states} - for label in labels.values(): - graph.add_state(label) - + labels = {state: f"s{rank}" for rank, state in enumerate(sorted(keep))} + for state in sorted(keep): + graph.add_state(labels[state]) used_inputs: set[Any] = set() used_outputs: set[Any] = set() - for state_id, histories in states.items(): - source = labels[state_id] + for state in sorted(keep): + longest = max(len(h) for h in states[state]) + aggregate = _state_aggregate(counts, {h for h in states[state] if len(h) == longest}) for input_symbol in counts.input_alphabet: - morph = counts.state_morph(histories, input_symbol) - if not morph: - continue - row: list[tuple[str, Any, float]] = [] - for output_symbol, prob in morph.items(): - if prob <= 0.0: - continue - emitting = [ - history for history in histories if _history_emits(counts, history, (input_symbol, output_symbol)) - ] - targets = {history_to_state.get((*history, (input_symbol, output_symbol))) for history in emitting} - targets.discard(None) - if len(targets) != 1: - continue - target_id = next(iter(targets)) - if target_id not in labels: - continue - row.append((labels[target_id], output_symbol, prob)) - total = sum(prob for _label, _out, prob in row) - if total <= 0.0: - continue - for target_label, output_symbol, prob in row: + row = [ + (labels[next(iter(edges[state][(input_symbol, output_symbol)]))], output_symbol, float(count)) + for output_symbol, count in sorted( + aggregate.get(input_symbol, Counter()).items(), key=lambda i: repr(i[0]) + ) + if count > 0 and edges[state].get((input_symbol, output_symbol)) + ] + total = sum(weight for _label, _output, weight in row) + for target_label, output_symbol, weight in row: graph.add_transition( - source, + labels[state], target_label, - **{ATTR_SYMBOL: input_symbol, ATTR_OUTPUT: output_symbol, ATTR_PROB: prob / total}, + **{ATTR_SYMBOL: input_symbol, ATTR_OUTPUT: output_symbol, ATTR_PROB: weight / total}, ) used_inputs.add(input_symbol) used_outputs.add(output_symbol) - total_visits = float(sum(visits.values())) - initial = {labels[state_id]: visits[state_id] / total_visits for state_id in states if visits.get(state_id, 0) > 0} - if not initial: - initial = {labels[next(iter(states))]: 1.0} - + kept_visits = {state: visits[state] for state in keep if visits[state] > 0} + total_visits = float(sum(kept_visits.values())) + initial = ( + {labels[state]: count / total_visits for state, count in kept_visits.items()} + if total_visits > 0 + else {labels[min(keep)]: 1.0} + ) result = EpsilonTransducer( input_alphabet=frozenset(used_inputs), output_alphabet=frozenset(used_outputs), @@ -491,9 +494,9 @@ def transcssr( max_length=max_length + 1, ) - states, history_to_state = _homogenize(counts, Lmax=max_length, alpha=alpha, test=test, min_count=min_count) - states = _determinize(states, history_to_state, counts) - history_to_state = {history: state_id for state_id, histories in states.items() for history in histories} - states = _drop_transient(states, history_to_state, counts) - history_to_state = {history: state_id for state_id, histories in states.items() for history in histories} - return _build_transducer(states, counts, history_to_state, xs, ys, length=max_length) + states = _homogenize(counts, Lmax=max_length, alpha=alpha, test=test, min_count=min_count) + everything = set(range(len(states))) + edges = _edges(states, counts, everything, Lmax=max_length, alpha=alpha, test=test) + alive = set().union(*_closed_classes(edges)) or everything + states, alive = _determinize(states, counts, alive, Lmax=max_length, alpha=alpha, test=test) + return _build_transducer(states, counts, alive, xs, ys, Lmax=max_length, alpha=alpha, test=test) diff --git a/sofic/generators/stack_hmm.py b/sofic/generators/stack_hmm.py index 37ea6cc..85f0b33 100644 --- a/sofic/generators/stack_hmm.py +++ b/sofic/generators/stack_hmm.py @@ -192,6 +192,9 @@ def word_probability(self, word: Sequence[Any]) -> float: current: dict[Configuration, float] = { (state, ()): float(prob) for state, prob in self.initial_distribution.items() if prob > _TOL } + # Masses are renormalized each step and pruned relative to their total; an absolute + # threshold would zero out every word longer than about 50 binary symbols. + log_scale = 0.0 for symbol in word: next_masses: dict[Configuration, float] = {} for config, mass in current.items(): @@ -199,10 +202,12 @@ def word_probability(self, word: Sequence[Any]) -> float: if transition.data.get(ATTR_SYMBOL) != symbol: continue next_masses[next_config] = next_masses.get(next_config, 0.0) + mass * prob - current = {config: mass for config, mass in next_masses.items() if mass > _TOL} - if not current: + total = sum(next_masses.values()) + if total <= 0.0: return 0.0 - return float(sum(current.values())) + log_scale += float(np.log(total)) + current = {config: mass / total for config, mass in next_masses.items() if mass > _TOL * total} + return float(np.exp(log_scale) * sum(current.values())) def words_of_length(self, length: int) -> dict[tuple[Any, ...], float]: """Return emitted words of ``length`` and their probabilities.""" diff --git a/sofic/generators/stack_inference.py b/sofic/generators/stack_inference.py index 8a6173a..84bc97f 100644 --- a/sofic/generators/stack_inference.py +++ b/sofic/generators/stack_inference.py @@ -12,11 +12,12 @@ History, SuffixCounts, _cluster_histories_by_morph, + _cssr_default_lmax, _cssr_determinize, - _cssr_homogenize, - _default_lmax, _drop_transient_states, _merge_similar_states, + morph_test_score, + morphs_differ, ) from sofic.generators.stack_hmm import HiddenMarkovStackModel from sofic.graph import ATTR_SYMBOL @@ -136,6 +137,29 @@ def successor(history: ConfigurationHistory, symbol: Any) -> ConfigurationHistor return successor +_RETURN = object() + + +def _control_counts(counts: StackSuffixCounts, alphabet: DyckAlphabet) -> StackSuffixCounts: + """Counts with every return symbol collapsed into one event. + + Which return symbol can follow is decided by the stack top (through matched + call-return pairs), not by the finite control, so comparing raw morphs would + split every control state by its stack top. + """ + returns = alphabet.return_alphabet + collapsed = StackSuffixCounts( + alphabet=tuple(symbol for symbol in counts.alphabet if symbol not in returns) + (_RETURN,) + ) + collapsed.history_counts = counts.history_counts + for history, nxt in counts.next_counts.items(): + merged: Counter[Any] = Counter() + for symbol, count in nxt.items(): + merged[_RETURN if symbol in returns else symbol] += count + collapsed.next_counts[history] = merged + return collapsed + + def _stack_homogenize( counts: StackSuffixCounts, *, @@ -144,14 +168,52 @@ def _stack_homogenize( alpha: float, test: Literal["g", "chi2", "tv"], max_stack_depth: int, + min_count: int = 1, ) -> tuple[dict[int, set[ConfigurationHistory]], dict[ConfigurationHistory, int]]: - return _cssr_homogenize( - counts, - Lmax=Lmax, - alpha=alpha, - test=test, - successor_fn=_stack_successor_fn(alphabet=alphabet, length=Lmax, max_stack_depth=max_stack_depth), - ) + """CSSR homogenization over ``(suffix, stack)`` configurations. + + Every observed stack contributes a root ``((), stack)``; suffixes then grow one + symbol into the past with their stack fixed, exactly as in flat CSSR. Growing + forward from the empty configuration instead only reaches stacks of depth at + most ``Lmax``, so deeper configurations had no state and their transitions were + dropped. Morphs are compared with return symbols collapsed (see :func:`_control_counts`). + """ + control = _control_counts(counts, alphabet) + states: dict[int, set[ConfigurationHistory]] = {0: {counts.empty_history}} + history_to_state: dict[ConfigurationHistory, int] = {counts.empty_history: 0} + + def place(child: ConfigurationHistory, parent_id: int) -> None: + target = parent_id + if morphs_differ(control, states[parent_id], {child}, alpha=alpha, test=test): + best_id, best_score = None, float("inf") + for candidate_id, candidate in states.items(): + if candidate_id == parent_id or morphs_differ(control, candidate, {child}, alpha=alpha, test=test): + continue + score = morph_test_score(control, candidate, {child}, test=test) + if score < best_score: + best_id, best_score = candidate_id, score + if best_id is None: + best_id = max(states) + 1 + states[best_id] = set() + target = best_id + states[target].add(child) + history_to_state[child] = target + + def observed(history: ConfigurationHistory) -> bool: + return sum(counts.next_counts.get(history, Counter()).values()) >= max(1, min_count) + + roots = {stack for suffix, stack in counts.history_counts if not suffix and stack} + for stack in sorted(roots, key=lambda stack: (len(stack), repr(stack))): + if observed(((), stack)): + place(((), stack), 0) + for length in range(Lmax): + for state_id in sorted(states): + for suffix, stack in sorted((h for h in states[state_id] if len(h[0]) == length), key=repr): + for symbol in counts.alphabet: + child = ((symbol, *suffix), stack) + if child not in history_to_state and observed(child): + place(child, state_id) + return states, history_to_state def _stack_determinize( @@ -180,8 +242,9 @@ def _stack_merge( *, alpha: float, test: Literal["g", "chi2", "tv"], + alphabet: DyckAlphabet, ) -> dict[int, set[ConfigurationHistory]]: - proxy = counts.restricted_to(set(history_to_state)) + proxy = _control_counts(counts, alphabet).restricted_to(set(history_to_state)) return _merge_similar_states(states, history_to_state, proxy, alpha=alpha, test=test) @@ -195,14 +258,13 @@ def _stack_drop_transient( max_stack_depth: int, ) -> dict[int, set[ConfigurationHistory]]: proxy = counts.restricted_to(set(history_to_state)) - return _drop_transient_states(states, history_to_state, proxy, length=length) - - -def _representative_stack(histories: set[ConfigurationHistory]) -> tuple[Any, ...]: - stacks = [stack for _suffix, stack in histories if stack] - if not stacks: - return () - return max(stacks, key=len) + return _drop_transient_states( + states, + history_to_state, + proxy, + length=length, + successor_fn=_stack_successor_fn(alphabet=alphabet, length=length, max_stack_depth=max_stack_depth), + ) def _counts_to_stack_hmm( @@ -213,18 +275,22 @@ def _counts_to_stack_hmm( *, alphabet: DyckAlphabet, length: int, + max_stack_depth: int, ) -> HiddenMarkovStackModel: visits: Counter[int] = Counter() seq = tuple(sequence) stack: list[Any] = [] for t in range(len(seq)): - for hist_len in range(0, min(t, length) + 1): - suffix = seq[t - hist_len : t] - state = history_to_state.get((suffix, tuple(stack))) + # Each step occupies one state: the one keyed by its longest available suffix. + for hist_len in range(min(t, length), -1, -1): + state = history_to_state.get((seq[t - hist_len : t], tuple(stack))) if state is not None: visits[state] += 1 + break symbol = seq[t] if symbol in alphabet.call_alphabet: + if len(stack) >= max_stack_depth: + stack = stack[1:] stack.append(symbol) elif symbol in alphabet.return_alphabet and stack: stack.pop() @@ -241,90 +307,64 @@ def _counts_to_stack_hmm( for label in state_labels.values(): model.graph.add_state(label) - call_refs: dict[tuple[Hashable, Any], TransitionRef] = {} - return_refs: dict[tuple[Hashable, Any, Any], TransitionRef] = {} + # Matched call-return pairs, as observed: a return ``r`` emitted with ``c`` on top. + observed_pairs = { + (stack[-1], symbol) + for (_suffix, stack), nxt in counts.next_counts.items() + if stack + for symbol, count in nxt.items() + if count > 0 and symbol in alphabet.return_alphabet + } + + def legal(symbol: Any, stack: tuple[Any, ...]) -> bool: + if symbol not in alphabet.return_alphabet: + return True + return not stack or (stack[-1], symbol) in observed_pairs + call_refs: dict[Any, list[TransitionRef]] = defaultdict(list) + return_refs: dict[Any, list[TransitionRef]] = defaultdict(list) for state_id, histories in states.items(): source = state_labels[state_id] - stack_repr = _representative_stack(histories) - stack_tops = {stack[-1] for _suffix, stack in histories if stack} - morph = counts.state_morph(histories) + emitted: Counter[Any] = Counter() + for history in histories: + emitted.update(counts.next_counts.get(history, Counter())) for symbol in counts.alphabet: - prob = morph[symbol] - if prob <= 0.0: + if emitted[symbol] <= 0: continue - emitting = [ - history for history in histories if counts.next_counts.get(history, Counter()).get(symbol, 0) > 0 - ] - if not emitting: - continue - child_histories = { - _successor_history( - history, - symbol, - alphabet=alphabet, - length=length, - max_stack_depth=max(len(stack_repr), 1), + targets: Counter[int] = Counter() + for history in histories: + count = counts.next_counts.get(history, Counter()).get(symbol, 0) + if count <= 0: + continue + child = _successor_history( + history, symbol, alphabet=alphabet, length=length, max_stack_depth=max_stack_depth ) - for history in emitting - } - targets = {history_to_state.get(child) for child in child_histories} - targets.discard(None) + target_id = history_to_state.get(child) + if target_id is not None: + targets[target_id] += count if not targets: continue - if len(targets) > 1: - target_counts: Counter[int] = Counter() - for history in emitting: - child = _successor_history( - history, - symbol, - alphabet=alphabet, - length=length, - max_stack_depth=max(len(stack_repr), 1), - ) - target_id = history_to_state.get(child) - if target_id is not None: - target_counts[target_id] += counts.history_counts.get(history, 0) - target_id = target_counts.most_common(1)[0][0] - else: - target_id = next(iter(targets)) - target = state_labels[target_id] - + target = state_labels[targets.most_common(1)[0][0]] + # The stack model renormalizes over the moves legal in each configuration, so + # a symbol's weight is its frequency among the visits where it was legal. + opportunities = sum( + sum(counts.next_counts.get(history, Counter()).values()) + for history in histories + if legal(symbol, history[1]) + ) + prob = emitted[symbol] / opportunities if symbol in alphabet.call_alphabet: - key = (source, symbol, target) - if key not in call_refs: - call_refs[key] = model.add_call_transition(source, target, symbol, prob) + call_refs[symbol].append(model.add_call_transition(source, target, symbol, prob)) elif symbol in alphabet.return_alphabet: - call_candidates = stack_tops or frozenset(alphabet.call_alphabet) - for matched_call in call_candidates: - key = (source, symbol, matched_call) - if key not in return_refs: - return_refs[key] = model.add_return_transition(source, target, symbol, prob) + return_refs[symbol].append(model.add_return_transition(source, target, symbol, prob)) else: model.add_internal_transition(source, target, symbol, prob) - for (_src, call_symbol, _target), call_ref in call_refs.items(): - for (_ret_source, _return_symbol, matched_call), return_ref in return_refs.items(): - if matched_call == call_symbol: + for call_symbol, return_symbol in observed_pairs: + for call_ref in call_refs.get(call_symbol, ()): + for return_ref in return_refs.get(return_symbol, ()): model.add_matched_pair(call_ref, return_ref) - for state_id, histories in states.items(): - source = state_labels[state_id] - for history in histories: - _suffix, stack = history - if not stack: - continue - for symbol in alphabet.return_alphabet: - if counts.next_counts.get(history, Counter()).get(symbol, 0) <= 0: - continue - matched_call = stack[-1] - for (src, sym, _tgt), call_ref in call_refs.items(): - if src != source or sym != matched_call: - continue - for (rsrc, rsym, mc), return_ref in return_refs.items(): - if rsrc == source and rsym == symbol and mc == matched_call: - model.add_matched_pair(call_ref, return_ref) - total_visits = float(sum(visits.values())) initial = {state_labels[state_id]: visits[state_id] / total_visits for state_id in states if visits[state_id] > 0} if not initial: @@ -348,7 +388,7 @@ def stack_cssr( seq = tuple(sequence) if len(seq) < 2: raise ValueError("sequence must contain at least two symbols") - max_length = Lmax if Lmax is not None else _default_lmax(len(seq), len(alphabet.symbol_alphabet), min_count) + max_length = Lmax if Lmax is not None else _cssr_default_lmax(len(seq), len(alphabet.symbol_alphabet)) counts = StackSuffixCounts.from_sequence( seq, alphabet=alphabet, @@ -362,6 +402,7 @@ def stack_cssr( alpha=alpha, test=test, max_stack_depth=max_stack_depth, + min_count=min_count, ) states = _stack_determinize( states, @@ -371,7 +412,7 @@ def stack_cssr( alphabet=alphabet, max_stack_depth=max_stack_depth, ) - states = _stack_merge(states, history_to_state, counts, alpha=alpha, test=test) + states = _stack_merge(states, history_to_state, counts, alpha=alpha, test=test, alphabet=alphabet) states = _stack_drop_transient( states, history_to_state, counts, length=max_length, alphabet=alphabet, max_stack_depth=max_stack_depth ) @@ -383,6 +424,7 @@ def stack_cssr( seq, alphabet=alphabet, length=max_length, + max_stack_depth=max_stack_depth, ) @@ -422,7 +464,7 @@ def stack_subtree_merge( history_to_state = { history: state_id for state_id, histories_in_state in states.items() for history in histories_in_state } - states = _stack_merge(states, history_to_state, counts, alpha=0.05, test="tv") + states = _stack_merge(states, history_to_state, counts, alpha=0.05, test="tv", alphabet=alphabet) history_to_state = { history: state_id for state_id, histories_in_state in states.items() for history in histories_in_state } @@ -439,6 +481,7 @@ def stack_subtree_merge( seq, alphabet=alphabet, length=L, + max_stack_depth=max_stack_depth, ) diff --git a/tests/test_epsilon_inference.py b/tests/test_epsilon_inference.py index be31dc6..47a2773 100644 --- a/tests/test_epsilon_inference.py +++ b/tests/test_epsilon_inference.py @@ -259,3 +259,17 @@ def test_cssr_non_synchronizable_process_returns_valid_machine(): inferred = cssr(observations, Lmax=4, alpha=0.001) inferred.validate() assert inferred.entropy_rate() >= oracle.entropy_rate() - 0.02 + + +@pytest.mark.parametrize(("name", "L", "n_states"), [("Even", 3, 2), ("GoldenMean", 2, 2), ("RkGM", 5, 8)]) +def test_subtree_merge_default_delta_recovers_process(name: str, L: int, n_states: int): + """Regression: the default delta = 0 compared sampled morphs to within 1e-3 and + successors were never truncated, so this raised StochasticValidationError.""" + from sofic.examples import processes + + oracle = processes.RkGM(5, 3) if name == "RkGM" else getattr(processes, name)() + observations, _ = sample(oracle, 20000, np.random.default_rng(5)) + inferred = subtree_merge(observations, L=L) + inferred.validate() + assert len(list(inferred.states())) == n_states + assert inferred.entropy_rate() == pytest.approx(oracle.entropy_rate(), abs=0.02) diff --git a/tests/test_epsilon_transducer_inference.py b/tests/test_epsilon_transducer_inference.py index f534091..66094c9 100644 --- a/tests/test_epsilon_transducer_inference.py +++ b/tests/test_epsilon_transducer_inference.py @@ -84,3 +84,33 @@ def test_reconstruction_reproduces_conditional_law(): def test_rejects_mismatched_lengths(): with pytest.raises(ValueError): transcssr("010", "01") + + +def _held_out_bits_per_symbol(eps: EpsilonTransducer, xs, ys, burn: int = 20) -> float: + states = list(eps.states()) + index = {state: i for i, state in enumerate(states)} + belief = np.array([eps.initial_distribution.get(state, 0.0) for state in states]) + total = 0.0 + for t, (x, y) in enumerate(zip(xs, ys, strict=True)): + nxt = np.zeros(len(states)) + for transition in eps.transitions(): + if transition.data["symbol"] == x and transition.data["output"] == y: + nxt[index[transition.target]] += belief[index[transition.source]] * transition.data["prob"] + mass = nxt.sum() + if mass <= 0.0: + return float("inf") + if t >= burn: + total -= np.log2(mass) + belief = nxt / mass + return total / (len(xs) - burn) + + +def test_recovers_two_step_delay(): + """Regression: joint suffixes grew forward and successors were never truncated, + so Delay(2) gave 5-19 states that forbade valid input-output pairs.""" + xs, ys = _paired_samples(Delay(2), 10000, seed=0) + eps = transcssr(xs, ys, input_alphabet=("0", "1"), output_alphabet=("0", "1")) + eps.validate() + assert len(list(eps.states())) == 4 + test_xs, test_ys = _paired_samples(Delay(2), 3000, seed=1) + assert _held_out_bits_per_symbol(eps, test_xs, test_ys) == pytest.approx(0.0, abs=1e-9) diff --git a/tests/test_stack_hmm.py b/tests/test_stack_hmm.py index a208c5e..b54d502 100644 --- a/tests/test_stack_hmm.py +++ b/tests/test_stack_hmm.py @@ -144,3 +144,16 @@ def test_validate_rejects_negative_probability(): with pytest.raises(StochasticValidationError): model.validate() + + +def test_word_probability_of_long_word_does_not_underflow(): + """Regression: masses below an absolute 1e-15 were pruned, so any word longer + than about 50 symbols had probability exactly 0.""" + from sofic.examples.shifts import dyck_shift_order + from sofic.shifts.sofic_dyck import transition_ref + + shift = dyck_shift_order(1) + refs = [transition_ref(transition) for transition in shift.transitions()] + model = HiddenMarkovStackModel.from_sofic_dyck_shift(shift, dict.fromkeys(refs, 1 / len(refs))) + word, _ = model.sample(200, rng=np.random.default_rng(0)) + assert np.log2(model.word_probability(word)) == pytest.approx(-200.0) diff --git a/tests/test_stack_inference.py b/tests/test_stack_inference.py index 86f3c83..e87af43 100644 --- a/tests/test_stack_inference.py +++ b/tests/test_stack_inference.py @@ -219,3 +219,28 @@ def test_benchmark_passive_paths(method: str): assert list(inferred.transitions()) oracle_prob = oracle.word_probability(tuple(observations[:12])) assert oracle_prob > 0.0 + + +def _held_out_bits_per_symbol(model, word) -> float: + return -np.log2(model.word_probability(tuple(word))) / len(word) + + +@pytest.mark.parametrize("Lmax", [2, 3]) +def test_stack_cssr_matches_motzkin_likelihood(Lmax: int): + """Regression: homogenization only reached stacks of depth <= Lmax and return + edges were paired with unobserved calls, so held-out words got probability 0.""" + shift = motzkin_shift() + oracle = HiddenMarkovStackModel.from_sofic_dyck_shift(shift, _uniform_probabilities(shift)) + alphabet = DyckAlphabet( + call_alphabet=shift.call_alphabet, + return_alphabet=shift.return_alphabet, + internal_alphabet=shift.internal_alphabet, + ) + observations, _ = oracle.sample(5000, rng=np.random.default_rng(0)) + held_out, _ = oracle.sample(60, rng=np.random.default_rng(99)) + inferred = stack_cssr(observations, alphabet=alphabet, Lmax=Lmax, max_stack_depth=4, alpha=0.001) + inferred.validate() + assert len(list(inferred.states())) == 1 + assert _held_out_bits_per_symbol(inferred, held_out) == pytest.approx( + _held_out_bits_per_symbol(oracle, held_out), abs=0.05 + )