sofic 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- sofic/__init__.py +185 -0
- sofic/automata/__init__.py +207 -0
- sofic/automata/_config_simulation.py +40 -0
- sofic/automata/active.py +611 -0
- sofic/automata/alergia.py +222 -0
- sofic/automata/algorithms.py +376 -0
- sofic/automata/atomaton.py +58 -0
- sofic/automata/base.py +161 -0
- sofic/automata/buchi.py +23 -0
- sofic/automata/buchi_simulation.py +67 -0
- sofic/automata/canonical_dual.py +18 -0
- sofic/automata/canonical_extraction.py +122 -0
- sofic/automata/dfa.py +85 -0
- sofic/automata/dfasat.py +195 -0
- sofic/automata/edsm.py +219 -0
- sofic/automata/enumeration.py +44 -0
- sofic/automata/icdfa.py +421 -0
- sofic/automata/idfa.py +363 -0
- sofic/automata/languages/__init__.py +39 -0
- sofic/automata/languages/_quotient_utils.py +64 -0
- sofic/automata/languages/atoms.py +31 -0
- sofic/automata/languages/automaton_ops.py +243 -0
- sofic/automata/languages/base.py +67 -0
- sofic/automata/languages/operations.py +78 -0
- sofic/automata/languages/quotients.py +66 -0
- sofic/automata/languages/residuals.py +25 -0
- sofic/automata/learning.py +79 -0
- sofic/automata/nfa.py +39 -0
- sofic/automata/nwa.py +343 -0
- sofic/automata/nwa_simulation.py +56 -0
- sofic/automata/observation.py +40 -0
- sofic/automata/papni.py +301 -0
- sofic/automata/regex.py +128 -0
- sofic/automata/rfsa.py +35 -0
- sofic/automata/rpni.py +193 -0
- sofic/automata/subsequential.py +201 -0
- sofic/automata/transducer_operations.py +350 -0
- sofic/automata/transducer_simulation.py +150 -0
- sofic/automata/transducers.py +365 -0
- sofic/automata/unifilar.py +107 -0
- sofic/automata/vpa.py +1373 -0
- sofic/automata/vpa_simulation.py +53 -0
- sofic/base.py +153 -0
- sofic/core.py +47 -0
- sofic/examples/__init__.py +86 -0
- sofic/examples/epsilon_machines.py +1089 -0
- sofic/examples/processes.py +1491 -0
- sofic/examples/shifts.py +144 -0
- sofic/exceptions.py +33 -0
- sofic/generators/__init__.py +115 -0
- sofic/generators/_word_measures.py +94 -0
- sofic/generators/alternative_complexity.py +104 -0
- sofic/generators/base.py +327 -0
- sofic/generators/bidirectional_construction.py +717 -0
- sofic/generators/bidirectional_epsilon_machine.py +689 -0
- sofic/generators/block_convergence.py +668 -0
- sofic/generators/block_entropy.py +578 -0
- sofic/generators/channel_measures.py +75 -0
- sofic/generators/conversions.py +182 -0
- sofic/generators/directional_flow.py +245 -0
- sofic/generators/edge_emissions.py +36 -0
- sofic/generators/edge_machine.py +178 -0
- sofic/generators/epsilon_construction.py +193 -0
- sofic/generators/epsilon_inference.py +703 -0
- sofic/generators/epsilon_machine.py +557 -0
- sofic/generators/epsilon_transducer.py +168 -0
- sofic/generators/epsilon_transducer_construction.py +185 -0
- sofic/generators/epsilon_transducer_inference.py +499 -0
- sofic/generators/hmm_inference.py +719 -0
- sofic/generators/information_diagram.py +428 -0
- sofic/generators/lumping.py +447 -0
- sofic/generators/markov.py +100 -0
- sofic/generators/mealy.py +156 -0
- sofic/generators/measures.py +257 -0
- sofic/generators/minimal_generative_model.py +821 -0
- sofic/generators/mixed_state.py +250 -0
- sofic/generators/mixed_state_construction.py +163 -0
- sofic/generators/moore.py +75 -0
- sofic/generators/nmachine.py +78 -0
- sofic/generators/nmachine_construction.py +70 -0
- sofic/generators/pfa.py +100 -0
- sofic/generators/prob.py +291 -0
- sofic/generators/process_equivalence.py +207 -0
- sofic/generators/quasi_inference.py +74 -0
- sofic/generators/quasi_realization.py +97 -0
- sofic/generators/reversal.py +66 -0
- sofic/generators/stack_hmm.py +426 -0
- sofic/generators/stack_inference.py +509 -0
- sofic/generators/stationary.py +134 -0
- sofic/generators/stochastic.py +65 -0
- sofic/generators/synchronization.py +407 -0
- sofic/generators/topological_epsilon_enumeration.py +349 -0
- sofic/generators/words.py +226 -0
- sofic/graph.py +135 -0
- sofic/indexing.py +31 -0
- sofic/inference/__init__.py +45 -0
- sofic/inference/bayesian/__init__.py +68 -0
- sofic/inference/bayesian/comparison.py +199 -0
- sofic/inference/bayesian/counts.py +219 -0
- sofic/inference/bayesian/diversity.py +254 -0
- sofic/inference/bayesian/epsilon.py +270 -0
- sofic/inference/bayesian/hdp_hmm.py +340 -0
- sofic/inference/bayesian/markov.py +294 -0
- sofic/inference/bayesian/pymc_backend.py +71 -0
- sofic/inference/bayesian/stack_hmm.py +215 -0
- sofic/inference/model_selection.py +365 -0
- sofic/inference/spectral.py +564 -0
- sofic/operations.py +16 -0
- sofic/properties.py +339 -0
- sofic/serialization.py +450 -0
- sofic/shifts/__init__.py +48 -0
- sofic/shifts/algorithms.py +84 -0
- sofic/shifts/base.py +49 -0
- sofic/shifts/cover_construction.py +76 -0
- sofic/shifts/covers.py +47 -0
- sofic/shifts/dyck_algorithms.py +100 -0
- sofic/shifts/dyck_enumeration.py +275 -0
- sofic/shifts/markov_dyck.py +172 -0
- sofic/shifts/parry_construction.py +82 -0
- sofic/shifts/sft.py +104 -0
- sofic/shifts/sft_construction.py +52 -0
- sofic/shifts/sliding_block_code.py +156 -0
- sofic/shifts/sofic.py +111 -0
- sofic/shifts/sofic_dyck.py +110 -0
- sofic/shifts/sofic_relation.py +64 -0
- sofic/shifts/textile.py +104 -0
- sofic/shifts/tmc.py +46 -0
- sofic/shifts/tmc_construction.py +58 -0
- sofic/shifts/topological_anatomy.py +150 -0
- sofic/states.py +27 -0
- sofic/testing/__init__.py +8 -0
- sofic/testing/strategies.py +154 -0
- sofic/viz/__init__.py +16 -0
- sofic/viz/_context.py +345 -0
- sofic/viz/_edge.py +216 -0
- sofic/viz/_format.py +89 -0
- sofic/viz/_labels.py +34 -0
- sofic/viz/_names.py +17 -0
- sofic/viz/_rational.py +20 -0
- sofic/viz/_tikz_compile.py +177 -0
- sofic/viz/_tikz_format.py +122 -0
- sofic/viz/_tikz_layout.py +218 -0
- sofic/viz/assets/vaucanson.tikz +71 -0
- sofic/viz/graphviz.py +158 -0
- sofic/viz/idiagram.py +350 -0
- sofic/viz/tikz.py +381 -0
- sofic-0.1.0.dist-info/METADATA +444 -0
- sofic-0.1.0.dist-info/RECORD +150 -0
- sofic-0.1.0.dist-info/WHEEL +4 -0
- sofic-0.1.0.dist-info/licenses/LICENSE.txt +29 -0
|
@@ -0,0 +1,821 @@
|
|
|
1
|
+
"""Minimal generative models from bidirectional epsilon-machines."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections import defaultdict
|
|
6
|
+
from collections.abc import Hashable
|
|
7
|
+
from typing import TYPE_CHECKING, Any, cast
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
|
|
11
|
+
from sofic.exceptions import StochasticValidationError
|
|
12
|
+
from sofic.generators.mealy import MealyHMM
|
|
13
|
+
from sofic.graph import ATTR_EMISSION, ATTR_PROB, TransitionGraph
|
|
14
|
+
|
|
15
|
+
if TYPE_CHECKING:
|
|
16
|
+
from sofic.generators.bidirectional_epsilon_machine import BidirectionalEpsilonMachine
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class _CommonInformationGenerativeModel(MealyHMM):
|
|
20
|
+
pair_state_channel: dict[tuple[Hashable, Hashable], dict[Hashable, float]]
|
|
21
|
+
joint_pair_distribution: dict[tuple[Hashable, Hashable], float]
|
|
22
|
+
source_bidirectional: BidirectionalEpsilonMachine | None
|
|
23
|
+
_entropy_rate: float | None
|
|
24
|
+
|
|
25
|
+
def __init__(
|
|
26
|
+
self,
|
|
27
|
+
*,
|
|
28
|
+
pair_state_channel: dict[tuple[Hashable, Hashable], dict[Hashable, float]] | None = None,
|
|
29
|
+
joint_pair_distribution: dict[tuple[Hashable, Hashable], float] | None = None,
|
|
30
|
+
source_bidirectional: BidirectionalEpsilonMachine | None = None,
|
|
31
|
+
entropy_rate: float | None = None,
|
|
32
|
+
**kwargs: Any,
|
|
33
|
+
) -> None:
|
|
34
|
+
super().__init__(**kwargs)
|
|
35
|
+
self.pair_state_channel = pair_state_channel or {}
|
|
36
|
+
self.joint_pair_distribution = joint_pair_distribution or {}
|
|
37
|
+
self.source_bidirectional = source_bidirectional
|
|
38
|
+
self._entropy_rate = entropy_rate
|
|
39
|
+
|
|
40
|
+
def entropy_rate(self) -> float:
|
|
41
|
+
"""Return the source process entropy rate when the source is known."""
|
|
42
|
+
if self._entropy_rate is not None:
|
|
43
|
+
return float(self._entropy_rate)
|
|
44
|
+
return super().entropy_rate()
|
|
45
|
+
|
|
46
|
+
def generative_complexity(self) -> float:
|
|
47
|
+
"""State entropy ``C_g`` of the generator."""
|
|
48
|
+
return self.state_entropy()
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
class MinimalGenerativeModel(_CommonInformationGenerativeModel):
|
|
52
|
+
"""Non-unifilar minimal-state-entropy generator.
|
|
53
|
+
|
|
54
|
+
The states are the optimized exact-common-information auxiliary variable
|
|
55
|
+
between the forward and reverse causal states of a bidirectional
|
|
56
|
+
epsilon-machine.
|
|
57
|
+
"""
|
|
58
|
+
|
|
59
|
+
exact_common_information: float
|
|
60
|
+
|
|
61
|
+
def __init__(self, *, exact_common_information: float = 0.0, **kwargs: Any) -> None:
|
|
62
|
+
super().__init__(**kwargs)
|
|
63
|
+
self.exact_common_information = float(exact_common_information)
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
class WynerGenerativeModel(_CommonInformationGenerativeModel):
|
|
67
|
+
"""Non-unifilar generator from a Wyner-common-information auxiliary.
|
|
68
|
+
|
|
69
|
+
The optimized value is ``I[(S+, S-) : G]`` and is stored as
|
|
70
|
+
``wyner_common_information``. The model's state entropy is ``H[G]`` and can
|
|
71
|
+
be strictly larger.
|
|
72
|
+
"""
|
|
73
|
+
|
|
74
|
+
wyner_common_information: float
|
|
75
|
+
|
|
76
|
+
def __init__(self, *, wyner_common_information: float = 0.0, **kwargs: Any) -> None:
|
|
77
|
+
super().__init__(**kwargs)
|
|
78
|
+
self.wyner_common_information = float(wyner_common_information)
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
class FunctionalGenerativeModel(_CommonInformationGenerativeModel):
|
|
82
|
+
"""Deterministic generator from a functional-common-information auxiliary.
|
|
83
|
+
|
|
84
|
+
The auxiliary state ``G`` is a *deterministic function* of the joint causal
|
|
85
|
+
state ``(S+, S-)`` -- the smallest such function rendering ``S+`` and ``S-``
|
|
86
|
+
conditionally independent. The optimized value ``H[G]`` is the functional
|
|
87
|
+
common information and is stored as ``functional_common_information``. Because
|
|
88
|
+
``G`` is deterministic, the model's state entropy ``H[G]`` equals that value
|
|
89
|
+
exactly.
|
|
90
|
+
"""
|
|
91
|
+
|
|
92
|
+
functional_common_information: float
|
|
93
|
+
|
|
94
|
+
def __init__(self, *, functional_common_information: float = 0.0, **kwargs: Any) -> None:
|
|
95
|
+
super().__init__(**kwargs)
|
|
96
|
+
self.functional_common_information = float(functional_common_information)
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
class GacsKornerGenerativeModel(_CommonInformationGenerativeModel):
|
|
100
|
+
"""Generator from the Gács-Körner (deterministic meet) auxiliary.
|
|
101
|
+
|
|
102
|
+
The states are the meet ``S+ ⩘ S-`` of the forward and reverse causal
|
|
103
|
+
states: the largest random variable that is simultaneously a deterministic
|
|
104
|
+
function of both. Unlike the Exact and Wyner models this is combinatorial,
|
|
105
|
+
not variational — each joint state pair maps to exactly one generative
|
|
106
|
+
state (its connected component in the joint support graph). Its state
|
|
107
|
+
entropy therefore equals the Gács-Körner common information
|
|
108
|
+
``K[S+ : S-]`` exactly, and captures only the conserved "core" (phase /
|
|
109
|
+
ergodic-component structure), which is often trivial for mixing processes.
|
|
110
|
+
"""
|
|
111
|
+
|
|
112
|
+
gk_common_information: float
|
|
113
|
+
|
|
114
|
+
def __init__(self, *, gk_common_information: float = 0.0, **kwargs: Any) -> None:
|
|
115
|
+
super().__init__(**kwargs)
|
|
116
|
+
self.gk_common_information = float(gk_common_information)
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
def minimal_generative_model(
|
|
120
|
+
bidir: BidirectionalEpsilonMachine,
|
|
121
|
+
*,
|
|
122
|
+
bound: int | None = None,
|
|
123
|
+
niter: int | None = 10,
|
|
124
|
+
maxiter: int = 1000,
|
|
125
|
+
polish: float | bool = 1e-8,
|
|
126
|
+
backend: str = "numpy",
|
|
127
|
+
cutoff: float = 1e-10,
|
|
128
|
+
reproduction_atol: float = 1e-3,
|
|
129
|
+
rng: np.random.Generator | None = None,
|
|
130
|
+
) -> MinimalGenerativeModel:
|
|
131
|
+
"""Construct a minimal generative model from a bidirectional epsilon-machine.
|
|
132
|
+
|
|
133
|
+
Parameters mirror :class:`dit.multivariate.common_informations.ExactCommonInformation`.
|
|
134
|
+
The returned model is an edge-emitting HMM over the optimized generative
|
|
135
|
+
state ``G``.
|
|
136
|
+
|
|
137
|
+
The exact-common-information optimizer is stochastic and can converge to the
|
|
138
|
+
right objective value while returning a generative channel that does not
|
|
139
|
+
reproduce the source process. The result is therefore verified against the
|
|
140
|
+
source process (word probabilities up to ``reproduction_atol``); if it fails,
|
|
141
|
+
the deterministic functional-common-information realization -- which always
|
|
142
|
+
reproduces the process and renders ``S+`` and ``S-`` conditionally
|
|
143
|
+
independent -- is returned instead.
|
|
144
|
+
"""
|
|
145
|
+
if cutoff < 0.0:
|
|
146
|
+
raise ValueError("cutoff must be nonnegative")
|
|
147
|
+
|
|
148
|
+
joint = _normalized_joint_distribution(bidir, cutoff=cutoff)
|
|
149
|
+
pair_channel, exact_common_information = _auxiliary_state_channel(
|
|
150
|
+
joint,
|
|
151
|
+
optimizer=_optimize_exact_common_information,
|
|
152
|
+
optimizer_name="exact common information",
|
|
153
|
+
bound=bound,
|
|
154
|
+
niter=niter,
|
|
155
|
+
maxiter=maxiter,
|
|
156
|
+
polish=polish,
|
|
157
|
+
backend=backend,
|
|
158
|
+
cutoff=cutoff,
|
|
159
|
+
rng=rng,
|
|
160
|
+
)
|
|
161
|
+
model = _model_from_channel(
|
|
162
|
+
bidir,
|
|
163
|
+
joint,
|
|
164
|
+
pair_channel,
|
|
165
|
+
model_cls=MinimalGenerativeModel,
|
|
166
|
+
model_name="minimal generative model",
|
|
167
|
+
measure_kwargs={"exact_common_information": exact_common_information},
|
|
168
|
+
cutoff=cutoff,
|
|
169
|
+
)
|
|
170
|
+
return cast(
|
|
171
|
+
MinimalGenerativeModel,
|
|
172
|
+
_ensure_reproducing(
|
|
173
|
+
bidir,
|
|
174
|
+
joint,
|
|
175
|
+
model,
|
|
176
|
+
model_cls=MinimalGenerativeModel,
|
|
177
|
+
model_name="minimal generative model",
|
|
178
|
+
measure_kwargs={"exact_common_information": exact_common_information},
|
|
179
|
+
cutoff=cutoff,
|
|
180
|
+
reproduction_atol=reproduction_atol,
|
|
181
|
+
),
|
|
182
|
+
)
|
|
183
|
+
|
|
184
|
+
|
|
185
|
+
def wyner_generative_model(
|
|
186
|
+
bidir: BidirectionalEpsilonMachine,
|
|
187
|
+
*,
|
|
188
|
+
bound: int | None = None,
|
|
189
|
+
niter: int | None = 10,
|
|
190
|
+
maxiter: int = 1000,
|
|
191
|
+
polish: float | bool = 1e-8,
|
|
192
|
+
backend: str = "numpy",
|
|
193
|
+
cutoff: float = 1e-10,
|
|
194
|
+
reproduction_atol: float = 1e-3,
|
|
195
|
+
rng: np.random.Generator | None = None,
|
|
196
|
+
) -> WynerGenerativeModel:
|
|
197
|
+
"""Construct a Wyner generative model from a bidirectional epsilon-machine.
|
|
198
|
+
|
|
199
|
+
The auxiliary state ``G`` is optimized for Wyner common information, i.e.
|
|
200
|
+
it minimizes ``I[(S+, S-) : G]`` subject to rendering ``S+`` and ``S-``
|
|
201
|
+
conditionally independent. The returned model's state entropy ``H[G]`` is
|
|
202
|
+
not generally equal to that optimized mutual information.
|
|
203
|
+
|
|
204
|
+
Like :func:`minimal_generative_model`, the stochastic optimizer's generative
|
|
205
|
+
channel is verified against the source process (word probabilities up to
|
|
206
|
+
``reproduction_atol``); on failure the deterministic functional realization
|
|
207
|
+
is returned instead (its ``H[G]`` may exceed the reported Wyner value).
|
|
208
|
+
"""
|
|
209
|
+
if cutoff < 0.0:
|
|
210
|
+
raise ValueError("cutoff must be nonnegative")
|
|
211
|
+
|
|
212
|
+
joint = _normalized_joint_distribution(bidir, cutoff=cutoff)
|
|
213
|
+
pair_channel, wyner_common_information = _auxiliary_state_channel(
|
|
214
|
+
joint,
|
|
215
|
+
optimizer=_optimize_wyner_common_information,
|
|
216
|
+
optimizer_name="Wyner common information",
|
|
217
|
+
bound=bound,
|
|
218
|
+
niter=niter,
|
|
219
|
+
maxiter=maxiter,
|
|
220
|
+
polish=polish,
|
|
221
|
+
backend=backend,
|
|
222
|
+
cutoff=cutoff,
|
|
223
|
+
rng=rng,
|
|
224
|
+
)
|
|
225
|
+
model = _model_from_channel(
|
|
226
|
+
bidir,
|
|
227
|
+
joint,
|
|
228
|
+
pair_channel,
|
|
229
|
+
model_cls=WynerGenerativeModel,
|
|
230
|
+
model_name="Wyner generative model",
|
|
231
|
+
measure_kwargs={"wyner_common_information": wyner_common_information},
|
|
232
|
+
cutoff=cutoff,
|
|
233
|
+
)
|
|
234
|
+
return cast(
|
|
235
|
+
WynerGenerativeModel,
|
|
236
|
+
_ensure_reproducing(
|
|
237
|
+
bidir,
|
|
238
|
+
joint,
|
|
239
|
+
model,
|
|
240
|
+
model_cls=WynerGenerativeModel,
|
|
241
|
+
model_name="Wyner generative model",
|
|
242
|
+
measure_kwargs={"wyner_common_information": wyner_common_information},
|
|
243
|
+
cutoff=cutoff,
|
|
244
|
+
reproduction_atol=reproduction_atol,
|
|
245
|
+
),
|
|
246
|
+
)
|
|
247
|
+
|
|
248
|
+
|
|
249
|
+
def functional_generative_model(
|
|
250
|
+
bidir: BidirectionalEpsilonMachine,
|
|
251
|
+
*,
|
|
252
|
+
cutoff: float = 1e-10,
|
|
253
|
+
strategy: str = "auto",
|
|
254
|
+
) -> FunctionalGenerativeModel:
|
|
255
|
+
"""Construct a functional generative model from a bidirectional epsilon-machine.
|
|
256
|
+
|
|
257
|
+
The auxiliary state ``G`` is the smallest *deterministic function* of the
|
|
258
|
+
joint causal state ``(S+, S-)`` that renders ``S+`` and ``S-`` conditionally
|
|
259
|
+
independent. Its entropy ``H[G]`` is the functional common information,
|
|
260
|
+
exposed as ``functional_common_information``. Because ``G`` is deterministic,
|
|
261
|
+
the model's state entropy equals that value exactly.
|
|
262
|
+
|
|
263
|
+
Unlike :func:`minimal_generative_model` and :func:`wyner_generative_model`,
|
|
264
|
+
the functional auxiliary is found by an exact partition search rather than a
|
|
265
|
+
stochastic optimizer, so the optimizer controls (``bound``, ``niter``,
|
|
266
|
+
``maxiter``, ``polish``, ``backend``, ``rng``) do not apply.
|
|
267
|
+
"""
|
|
268
|
+
if cutoff < 0.0:
|
|
269
|
+
raise ValueError("cutoff must be nonnegative")
|
|
270
|
+
|
|
271
|
+
joint = _normalized_joint_distribution(bidir, cutoff=cutoff)
|
|
272
|
+
pair_channel, functional_common_information = _functional_state_channel(
|
|
273
|
+
joint,
|
|
274
|
+
cutoff=cutoff,
|
|
275
|
+
strategy=strategy,
|
|
276
|
+
)
|
|
277
|
+
return cast(
|
|
278
|
+
FunctionalGenerativeModel,
|
|
279
|
+
_model_from_channel(
|
|
280
|
+
bidir,
|
|
281
|
+
joint,
|
|
282
|
+
pair_channel,
|
|
283
|
+
model_cls=FunctionalGenerativeModel,
|
|
284
|
+
model_name="functional generative model",
|
|
285
|
+
measure_kwargs={"functional_common_information": functional_common_information},
|
|
286
|
+
cutoff=cutoff,
|
|
287
|
+
),
|
|
288
|
+
)
|
|
289
|
+
|
|
290
|
+
|
|
291
|
+
def gacs_korner_generative_model(
|
|
292
|
+
bidir: BidirectionalEpsilonMachine,
|
|
293
|
+
*,
|
|
294
|
+
cutoff: float = 1e-10,
|
|
295
|
+
) -> GacsKornerGenerativeModel:
|
|
296
|
+
"""Construct a Gács-Körner generative model from a bidirectional epsilon-machine.
|
|
297
|
+
|
|
298
|
+
The generative state ``G`` is the meet ``S+ ⩘ S-`` of the forward and
|
|
299
|
+
reverse causal states — the largest random variable that is simultaneously
|
|
300
|
+
a deterministic function of both, obtained combinatorially as the connected
|
|
301
|
+
components of the joint support graph. Because the meet is deterministic
|
|
302
|
+
there is nothing to optimize, so this factory takes no optimizer arguments;
|
|
303
|
+
the returned model's state entropy ``H[G]`` equals the Gács-Körner common
|
|
304
|
+
information ``K[S+ : S-]``.
|
|
305
|
+
"""
|
|
306
|
+
if cutoff < 0.0:
|
|
307
|
+
raise ValueError("cutoff must be nonnegative")
|
|
308
|
+
|
|
309
|
+
joint = _normalized_joint_distribution(bidir, cutoff=cutoff)
|
|
310
|
+
pair_channel = _gacs_korner_meet_channel(joint)
|
|
311
|
+
component_mass: dict[Hashable, float] = defaultdict(float)
|
|
312
|
+
for pair, pair_mass in joint.items():
|
|
313
|
+
(state,) = pair_channel[pair]
|
|
314
|
+
component_mass[state] += pair_mass
|
|
315
|
+
gk_common_information = _entropy(component_mass.values())
|
|
316
|
+
return cast(
|
|
317
|
+
GacsKornerGenerativeModel,
|
|
318
|
+
_model_from_channel(
|
|
319
|
+
bidir,
|
|
320
|
+
joint,
|
|
321
|
+
pair_channel,
|
|
322
|
+
model_cls=GacsKornerGenerativeModel,
|
|
323
|
+
model_name="Gács-Körner generative model",
|
|
324
|
+
measure_kwargs={"gk_common_information": gk_common_information},
|
|
325
|
+
cutoff=cutoff,
|
|
326
|
+
),
|
|
327
|
+
)
|
|
328
|
+
|
|
329
|
+
|
|
330
|
+
def _gacs_korner_meet_channel(
|
|
331
|
+
joint: dict[tuple[Hashable, Hashable], float],
|
|
332
|
+
) -> dict[tuple[Hashable, Hashable], dict[Hashable, float]]:
|
|
333
|
+
"""Deterministic channel assigning each state pair to its meet component.
|
|
334
|
+
|
|
335
|
+
Builds the bipartite graph linking a forward state ``alpha`` to a reverse
|
|
336
|
+
state ``gamma`` whenever ``p(alpha, gamma) > 0``; the connected components
|
|
337
|
+
are the atoms of the meet ``S+ ⩘ S-``.
|
|
338
|
+
"""
|
|
339
|
+
import networkx as nx
|
|
340
|
+
|
|
341
|
+
graph = nx.Graph()
|
|
342
|
+
for alpha, gamma in joint:
|
|
343
|
+
graph.add_edge(("+", alpha), ("-", gamma))
|
|
344
|
+
|
|
345
|
+
component_of: dict[Hashable, int] = {}
|
|
346
|
+
for rank, component in enumerate(nx.connected_components(graph)):
|
|
347
|
+
for node in component:
|
|
348
|
+
component_of[node] = rank
|
|
349
|
+
|
|
350
|
+
return {pair: {f"G{component_of[('+', pair[0])]}": 1.0} for pair in joint}
|
|
351
|
+
|
|
352
|
+
|
|
353
|
+
def _functional_state_channel(
|
|
354
|
+
joint: dict[tuple[Hashable, Hashable], float],
|
|
355
|
+
*,
|
|
356
|
+
cutoff: float,
|
|
357
|
+
strategy: str,
|
|
358
|
+
) -> tuple[dict[tuple[Hashable, Hashable], dict[Hashable, float]], float]:
|
|
359
|
+
"""Deterministic channel from the functional-common-information partition.
|
|
360
|
+
|
|
361
|
+
Runs dit's exact functional-Markov search on the joint ``(S+, S-)``
|
|
362
|
+
distribution, recovers the optimal outcome partition ``W = f(S+, S-)``, and
|
|
363
|
+
maps each state pair to its (single) block label. The returned value is the
|
|
364
|
+
functional common information ``H[W]``.
|
|
365
|
+
"""
|
|
366
|
+
plus_states = {pair[0] for pair in joint}
|
|
367
|
+
minus_states = {pair[1] for pair in joint}
|
|
368
|
+
if len(plus_states) <= 1 or len(minus_states) <= 1:
|
|
369
|
+
return {pair: {"G0": 1.0} for pair in joint}, 0.0
|
|
370
|
+
matching_channel = _matching_support_channel(joint)
|
|
371
|
+
if matching_channel is not None:
|
|
372
|
+
return matching_channel, _entropy(joint.values())
|
|
373
|
+
|
|
374
|
+
from dit.multivariate.common_informations.functional_common_information import functional_markov_chain
|
|
375
|
+
|
|
376
|
+
dist, plus_index, minus_index = _indexed_joint_distribution(joint)
|
|
377
|
+
|
|
378
|
+
stats: dict[str, Any] = {}
|
|
379
|
+
value = _as_float(functional_markov_chain(dist, [[0], [1]], _strategy=strategy, _stats=stats))
|
|
380
|
+
partition = stats.get("partition")
|
|
381
|
+
if partition is None:
|
|
382
|
+
raise StochasticValidationError("functional common information search returned no partition")
|
|
383
|
+
|
|
384
|
+
outcome_rank: dict[tuple[int, int], int] = {}
|
|
385
|
+
for rank, block in enumerate(partition):
|
|
386
|
+
for outcome in block:
|
|
387
|
+
outcome_rank[tuple(outcome)] = rank
|
|
388
|
+
|
|
389
|
+
pair_rank: dict[tuple[Hashable, Hashable], int] = {}
|
|
390
|
+
for pair in joint:
|
|
391
|
+
key = (plus_index[pair[0]], minus_index[pair[1]])
|
|
392
|
+
if key not in outcome_rank:
|
|
393
|
+
raise StochasticValidationError(f"functional partition is missing joint state {pair!r}")
|
|
394
|
+
pair_rank[pair] = outcome_rank[key]
|
|
395
|
+
|
|
396
|
+
used_ranks = sorted(set(pair_rank.values()))
|
|
397
|
+
relabel = {rank: f"G{i}" for i, rank in enumerate(used_ranks)}
|
|
398
|
+
channel = {pair: {relabel[rank]: 1.0} for pair, rank in pair_rank.items()}
|
|
399
|
+
return channel, value
|
|
400
|
+
|
|
401
|
+
|
|
402
|
+
def _require_dit():
|
|
403
|
+
from sofic.generators.measures import require_dit
|
|
404
|
+
|
|
405
|
+
return require_dit("minimal generative models")
|
|
406
|
+
|
|
407
|
+
|
|
408
|
+
def _as_numpy(array: Any) -> np.ndarray:
|
|
409
|
+
if hasattr(array, "detach"):
|
|
410
|
+
array = array.detach().cpu().numpy()
|
|
411
|
+
return np.asarray(array, dtype=float)
|
|
412
|
+
|
|
413
|
+
|
|
414
|
+
def _as_float(value: Any) -> float:
|
|
415
|
+
if hasattr(value, "detach"):
|
|
416
|
+
value = value.detach().cpu().item()
|
|
417
|
+
return float(value)
|
|
418
|
+
|
|
419
|
+
|
|
420
|
+
def _normalized_joint_distribution(
|
|
421
|
+
bidir: BidirectionalEpsilonMachine,
|
|
422
|
+
*,
|
|
423
|
+
cutoff: float,
|
|
424
|
+
) -> dict[tuple[Hashable, Hashable], float]:
|
|
425
|
+
joint = {pair: float(mass) for pair, mass in bidir.joint_distribution().items() if float(mass) > cutoff}
|
|
426
|
+
total = sum(joint.values())
|
|
427
|
+
if total <= 0.0:
|
|
428
|
+
raise StochasticValidationError("bidirectional machine has empty stationary joint distribution")
|
|
429
|
+
return {pair: mass / total for pair, mass in joint.items()}
|
|
430
|
+
|
|
431
|
+
|
|
432
|
+
def _indexed_joint_distribution(
|
|
433
|
+
joint: dict[tuple[Hashable, Hashable], float],
|
|
434
|
+
) -> tuple[Any, dict[Hashable, int], dict[Hashable, int]]:
|
|
435
|
+
dit = _require_dit()
|
|
436
|
+
plus_states = tuple(sorted({pair[0] for pair in joint}, key=repr))
|
|
437
|
+
minus_states = tuple(sorted({pair[1] for pair in joint}, key=repr))
|
|
438
|
+
plus_index = {state: i for i, state in enumerate(plus_states)}
|
|
439
|
+
minus_index = {state: i for i, state in enumerate(minus_states)}
|
|
440
|
+
|
|
441
|
+
outcomes = [(plus_index[alpha], minus_index[gamma]) for alpha, gamma in joint]
|
|
442
|
+
probs = [joint[pair] for pair in joint]
|
|
443
|
+
return dit.Distribution(outcomes, probs), plus_index, minus_index
|
|
444
|
+
|
|
445
|
+
|
|
446
|
+
def _optimized_auxiliary_joint(
|
|
447
|
+
dist: Any,
|
|
448
|
+
*,
|
|
449
|
+
optimizer: Any,
|
|
450
|
+
optimizer_name: str,
|
|
451
|
+
bound: int | None,
|
|
452
|
+
niter: int | None,
|
|
453
|
+
maxiter: int,
|
|
454
|
+
polish: float | bool,
|
|
455
|
+
backend: str,
|
|
456
|
+
rng: np.random.Generator | None,
|
|
457
|
+
) -> tuple[Any, np.ndarray]:
|
|
458
|
+
opt = optimizer(
|
|
459
|
+
dist,
|
|
460
|
+
bound=bound,
|
|
461
|
+
niter=niter,
|
|
462
|
+
maxiter=maxiter,
|
|
463
|
+
polish=polish,
|
|
464
|
+
backend=backend,
|
|
465
|
+
rng=rng,
|
|
466
|
+
)
|
|
467
|
+
|
|
468
|
+
aux_joint = _as_numpy(opt.construct_joint(opt._optima))
|
|
469
|
+
if aux_joint.ndim < 3:
|
|
470
|
+
raise StochasticValidationError(f"{optimizer_name} optimizer returned an invalid joint shape")
|
|
471
|
+
if aux_joint.ndim > 3:
|
|
472
|
+
aux_joint = aux_joint.sum(axis=tuple(range(2, aux_joint.ndim - 1)))
|
|
473
|
+
return opt, np.maximum(aux_joint, 0.0)
|
|
474
|
+
|
|
475
|
+
|
|
476
|
+
def _raw_channel_from_auxiliary_joint(
|
|
477
|
+
aux_joint: np.ndarray,
|
|
478
|
+
joint: dict[tuple[Hashable, Hashable], float],
|
|
479
|
+
plus_index: dict[Hashable, int],
|
|
480
|
+
minus_index: dict[Hashable, int],
|
|
481
|
+
*,
|
|
482
|
+
optimizer_name: str,
|
|
483
|
+
cutoff: float,
|
|
484
|
+
polish: float | bool,
|
|
485
|
+
niter: int | None,
|
|
486
|
+
) -> tuple[dict[tuple[Hashable, Hashable], np.ndarray] | None, str | None]:
|
|
487
|
+
raw_channel: dict[tuple[Hashable, Hashable], np.ndarray] = {}
|
|
488
|
+
for pair, target_mass in joint.items():
|
|
489
|
+
alpha, gamma = pair
|
|
490
|
+
row = aux_joint[plus_index[alpha], minus_index[gamma], :]
|
|
491
|
+
total = float(row.sum())
|
|
492
|
+
if total <= cutoff:
|
|
493
|
+
return None, (
|
|
494
|
+
f"missing generative-state channel row for joint state {pair!r}: "
|
|
495
|
+
f"target joint mass={target_mass:.17g}, returned row mass={total:.17g}, "
|
|
496
|
+
f"cutoff={cutoff:.17g}, polish={polish!r}, niter={niter!r}, optimizer={optimizer_name!r}"
|
|
497
|
+
)
|
|
498
|
+
raw_channel[pair] = row / total
|
|
499
|
+
return raw_channel, None
|
|
500
|
+
|
|
501
|
+
|
|
502
|
+
def _auxiliary_state_channel(
|
|
503
|
+
joint: dict[tuple[Hashable, Hashable], float],
|
|
504
|
+
*,
|
|
505
|
+
optimizer: Any,
|
|
506
|
+
optimizer_name: str,
|
|
507
|
+
bound: int | None,
|
|
508
|
+
niter: int | None,
|
|
509
|
+
maxiter: int,
|
|
510
|
+
polish: float | bool,
|
|
511
|
+
backend: str,
|
|
512
|
+
cutoff: float,
|
|
513
|
+
rng: np.random.Generator | None,
|
|
514
|
+
) -> tuple[dict[tuple[Hashable, Hashable], dict[Hashable, float]], float]:
|
|
515
|
+
plus_states = {pair[0] for pair in joint}
|
|
516
|
+
minus_states = {pair[1] for pair in joint}
|
|
517
|
+
if len(plus_states) <= 1 or len(minus_states) <= 1:
|
|
518
|
+
return {pair: {"G0": 1.0} for pair in joint}, 0.0
|
|
519
|
+
matching_channel = _matching_support_channel(joint)
|
|
520
|
+
if matching_channel is not None:
|
|
521
|
+
return matching_channel, _entropy(joint.values())
|
|
522
|
+
|
|
523
|
+
dist, plus_index, minus_index = _indexed_joint_distribution(joint)
|
|
524
|
+
opt, aux_joint = _optimized_auxiliary_joint(
|
|
525
|
+
dist,
|
|
526
|
+
optimizer=optimizer,
|
|
527
|
+
optimizer_name=optimizer_name,
|
|
528
|
+
bound=bound,
|
|
529
|
+
niter=niter,
|
|
530
|
+
maxiter=maxiter,
|
|
531
|
+
polish=polish,
|
|
532
|
+
backend=backend,
|
|
533
|
+
rng=rng,
|
|
534
|
+
)
|
|
535
|
+
raw_channel, validation_error = _raw_channel_from_auxiliary_joint(
|
|
536
|
+
aux_joint,
|
|
537
|
+
joint,
|
|
538
|
+
plus_index,
|
|
539
|
+
minus_index,
|
|
540
|
+
optimizer_name=optimizer_name,
|
|
541
|
+
cutoff=cutoff,
|
|
542
|
+
polish=polish,
|
|
543
|
+
niter=niter,
|
|
544
|
+
)
|
|
545
|
+
if validation_error is not None and polish:
|
|
546
|
+
opt, aux_joint = _optimized_auxiliary_joint(
|
|
547
|
+
dist,
|
|
548
|
+
optimizer=optimizer,
|
|
549
|
+
optimizer_name=optimizer_name,
|
|
550
|
+
bound=bound,
|
|
551
|
+
niter=niter,
|
|
552
|
+
maxiter=maxiter,
|
|
553
|
+
polish=False,
|
|
554
|
+
backend=backend,
|
|
555
|
+
rng=rng,
|
|
556
|
+
)
|
|
557
|
+
raw_channel, retry_error = _raw_channel_from_auxiliary_joint(
|
|
558
|
+
aux_joint,
|
|
559
|
+
joint,
|
|
560
|
+
plus_index,
|
|
561
|
+
minus_index,
|
|
562
|
+
optimizer_name=optimizer_name,
|
|
563
|
+
cutoff=cutoff,
|
|
564
|
+
polish=False,
|
|
565
|
+
niter=niter,
|
|
566
|
+
)
|
|
567
|
+
if retry_error is not None:
|
|
568
|
+
raise StochasticValidationError(f"{validation_error}; retry without polishing also failed: {retry_error}")
|
|
569
|
+
elif validation_error is not None:
|
|
570
|
+
raise StochasticValidationError(validation_error)
|
|
571
|
+
|
|
572
|
+
assert raw_channel is not None
|
|
573
|
+
state_masses = _state_masses(joint, raw_channel)
|
|
574
|
+
active_indices = [i for i, mass in enumerate(state_masses) if mass > cutoff]
|
|
575
|
+
if not active_indices:
|
|
576
|
+
active_indices = [int(np.argmax(state_masses))]
|
|
577
|
+
state_labels = {index: f"G{rank}" for rank, index in enumerate(active_indices)}
|
|
578
|
+
|
|
579
|
+
channel: dict[tuple[Hashable, Hashable], dict[Hashable, float]] = {}
|
|
580
|
+
for pair, row in raw_channel.items():
|
|
581
|
+
entries = {state_labels[i]: float(row[i]) for i in active_indices if row[i] > cutoff}
|
|
582
|
+
total = sum(entries.values())
|
|
583
|
+
if total <= cutoff:
|
|
584
|
+
best = max(active_indices, key=lambda i: row[i])
|
|
585
|
+
entries = {state_labels[best]: 1.0}
|
|
586
|
+
else:
|
|
587
|
+
entries = {state: prob / total for state, prob in entries.items()}
|
|
588
|
+
channel[pair] = entries
|
|
589
|
+
|
|
590
|
+
optimized_value = _as_float(opt.objective(opt._optima))
|
|
591
|
+
return channel, optimized_value
|
|
592
|
+
|
|
593
|
+
|
|
594
|
+
def _matching_support_channel(
|
|
595
|
+
joint: dict[tuple[Hashable, Hashable], float],
|
|
596
|
+
) -> dict[tuple[Hashable, Hashable], dict[Hashable, float]] | None:
|
|
597
|
+
plus_to_minus: dict[Hashable, Hashable] = {}
|
|
598
|
+
minus_to_plus: dict[Hashable, Hashable] = {}
|
|
599
|
+
for alpha, gamma in joint:
|
|
600
|
+
if alpha in plus_to_minus and plus_to_minus[alpha] != gamma:
|
|
601
|
+
return None
|
|
602
|
+
if gamma in minus_to_plus and minus_to_plus[gamma] != alpha:
|
|
603
|
+
return None
|
|
604
|
+
plus_to_minus[alpha] = gamma
|
|
605
|
+
minus_to_plus[gamma] = alpha
|
|
606
|
+
|
|
607
|
+
return {pair: {f"G{i}": 1.0} for i, pair in enumerate(sorted(joint, key=repr))}
|
|
608
|
+
|
|
609
|
+
|
|
610
|
+
def _entropy(masses: Any) -> float:
|
|
611
|
+
from sofic.generators.stochastic import shannon_entropy
|
|
612
|
+
|
|
613
|
+
return shannon_entropy(masses, atol=0.0)
|
|
614
|
+
|
|
615
|
+
|
|
616
|
+
def _optimize_exact_common_information(
|
|
617
|
+
dist: Any,
|
|
618
|
+
*,
|
|
619
|
+
bound: int | None,
|
|
620
|
+
niter: int | None,
|
|
621
|
+
maxiter: int,
|
|
622
|
+
polish: float | bool,
|
|
623
|
+
backend: str,
|
|
624
|
+
rng: np.random.Generator | None,
|
|
625
|
+
) -> Any:
|
|
626
|
+
from dit.multivariate._backend import _make_backend_subclass
|
|
627
|
+
from dit.multivariate.common_informations.exact_common_information import ExactCommonInformation
|
|
628
|
+
|
|
629
|
+
cls = _make_backend_subclass(ExactCommonInformation, backend)
|
|
630
|
+
opt = cls(dist, [[0], [1]], bound=bound)
|
|
631
|
+
if isinstance(polish, int | float) and not isinstance(polish, bool) and polish > 0.0:
|
|
632
|
+
options = opt._additional_options.setdefault("options", {})
|
|
633
|
+
options["ftol"] = min(float(polish), float(options.get("ftol", polish)))
|
|
634
|
+
opt.optimize(niter=niter, maxiter=maxiter, polish=polish, rng=rng)
|
|
635
|
+
return opt
|
|
636
|
+
|
|
637
|
+
|
|
638
|
+
def _optimize_wyner_common_information(
|
|
639
|
+
dist: Any,
|
|
640
|
+
*,
|
|
641
|
+
bound: int | None,
|
|
642
|
+
niter: int | None,
|
|
643
|
+
maxiter: int,
|
|
644
|
+
polish: float | bool,
|
|
645
|
+
backend: str,
|
|
646
|
+
rng: np.random.Generator | None,
|
|
647
|
+
) -> Any:
|
|
648
|
+
from dit.multivariate._backend import _make_backend_subclass
|
|
649
|
+
from dit.multivariate.common_informations.wyner_common_information import WynerCommonInformation
|
|
650
|
+
|
|
651
|
+
cls = _make_backend_subclass(WynerCommonInformation, backend)
|
|
652
|
+
opt = cls(dist, [[0], [1]], bound=bound)
|
|
653
|
+
if isinstance(polish, int | float) and not isinstance(polish, bool) and polish > 0.0:
|
|
654
|
+
options = opt._additional_options.setdefault("options", {})
|
|
655
|
+
options["ftol"] = min(float(polish), float(options.get("ftol", polish)))
|
|
656
|
+
opt.optimize(niter=niter, maxiter=maxiter, polish=polish, rng=rng)
|
|
657
|
+
return opt
|
|
658
|
+
|
|
659
|
+
|
|
660
|
+
def _state_masses(
|
|
661
|
+
joint: dict[tuple[Hashable, Hashable], float],
|
|
662
|
+
pair_channel: dict[tuple[Hashable, Hashable], np.ndarray],
|
|
663
|
+
) -> np.ndarray:
|
|
664
|
+
size = max(len(row) for row in pair_channel.values())
|
|
665
|
+
masses = np.zeros(size, dtype=float)
|
|
666
|
+
for pair, pair_mass in joint.items():
|
|
667
|
+
masses += pair_mass * pair_channel[pair]
|
|
668
|
+
return masses
|
|
669
|
+
|
|
670
|
+
|
|
671
|
+
def _reproduction_length(reference: Any) -> int:
|
|
672
|
+
"""Word length at which to compare a generative model against the source process."""
|
|
673
|
+
n_states = len(list(reference.states()))
|
|
674
|
+
return min(8, max(4, 2 * n_states))
|
|
675
|
+
|
|
676
|
+
|
|
677
|
+
def _reproduction_error(model: Any, reference: Any, *, max_length: int) -> float:
|
|
678
|
+
"""Max abs word-probability deviation of ``model`` from ``reference`` up to ``max_length``."""
|
|
679
|
+
error = 0.0
|
|
680
|
+
for length in range(max_length + 1):
|
|
681
|
+
model_words = model.word_probabilities(length, sparse=False)
|
|
682
|
+
reference_words = reference.word_probabilities(length, sparse=False)
|
|
683
|
+
for word in set(model_words) | set(reference_words):
|
|
684
|
+
error = max(error, abs(model_words.get(word, 0.0) - reference_words.get(word, 0.0)))
|
|
685
|
+
return error
|
|
686
|
+
|
|
687
|
+
|
|
688
|
+
def _ensure_reproducing(
|
|
689
|
+
bidir: BidirectionalEpsilonMachine,
|
|
690
|
+
joint: dict[tuple[Hashable, Hashable], float],
|
|
691
|
+
model: _CommonInformationGenerativeModel,
|
|
692
|
+
*,
|
|
693
|
+
model_cls: type[_CommonInformationGenerativeModel],
|
|
694
|
+
model_name: str,
|
|
695
|
+
measure_kwargs: dict[str, float],
|
|
696
|
+
cutoff: float,
|
|
697
|
+
reproduction_atol: float,
|
|
698
|
+
strategy: str = "auto",
|
|
699
|
+
) -> _CommonInformationGenerativeModel:
|
|
700
|
+
"""Return ``model`` if it reproduces the source process, else the functional fallback.
|
|
701
|
+
|
|
702
|
+
The exact-/Wyner-common-information optimizers are stochastic and can return a
|
|
703
|
+
channel that fails to reproduce the process even when the objective value is
|
|
704
|
+
correct. The deterministic functional-common-information channel always
|
|
705
|
+
reproduces the process and renders ``S+`` and ``S-`` conditionally
|
|
706
|
+
independent, so it is used as an exact fallback (the optimizer's reported
|
|
707
|
+
measure value is preserved). Raises if even the functional realization fails.
|
|
708
|
+
"""
|
|
709
|
+
reference = bidir.forward_machine
|
|
710
|
+
max_length = _reproduction_length(reference)
|
|
711
|
+
if _reproduction_error(model, reference, max_length=max_length) <= reproduction_atol:
|
|
712
|
+
return model
|
|
713
|
+
|
|
714
|
+
functional_channel, _functional_value = _functional_state_channel(joint, cutoff=cutoff, strategy=strategy)
|
|
715
|
+
fallback = _model_from_channel(
|
|
716
|
+
bidir,
|
|
717
|
+
joint,
|
|
718
|
+
functional_channel,
|
|
719
|
+
model_cls=model_cls,
|
|
720
|
+
model_name=model_name,
|
|
721
|
+
measure_kwargs=measure_kwargs,
|
|
722
|
+
cutoff=cutoff,
|
|
723
|
+
)
|
|
724
|
+
if _reproduction_error(fallback, reference, max_length=max_length) > reproduction_atol:
|
|
725
|
+
raise StochasticValidationError(
|
|
726
|
+
f"{model_name} does not reproduce the source process within {reproduction_atol:g} "
|
|
727
|
+
"and the deterministic functional fallback also failed"
|
|
728
|
+
)
|
|
729
|
+
return fallback
|
|
730
|
+
|
|
731
|
+
|
|
732
|
+
def _model_from_channel(
|
|
733
|
+
bidir: BidirectionalEpsilonMachine,
|
|
734
|
+
joint: dict[tuple[Hashable, Hashable], float],
|
|
735
|
+
pair_channel: dict[tuple[Hashable, Hashable], dict[Hashable, float]],
|
|
736
|
+
*,
|
|
737
|
+
model_cls: type[_CommonInformationGenerativeModel],
|
|
738
|
+
model_name: str,
|
|
739
|
+
measure_kwargs: dict[str, float],
|
|
740
|
+
cutoff: float,
|
|
741
|
+
) -> _CommonInformationGenerativeModel:
|
|
742
|
+
state_mass: dict[Hashable, float] = defaultdict(float)
|
|
743
|
+
for pair, pair_mass in joint.items():
|
|
744
|
+
for state, prob in pair_channel[pair].items():
|
|
745
|
+
state_mass[state] += pair_mass * prob
|
|
746
|
+
|
|
747
|
+
active_states = tuple(
|
|
748
|
+
state for state, mass in sorted(state_mass.items(), key=lambda item: repr(item[0])) if mass > cutoff
|
|
749
|
+
)
|
|
750
|
+
total_state_mass = sum(state_mass[state] for state in active_states)
|
|
751
|
+
if total_state_mass <= 0.0:
|
|
752
|
+
raise StochasticValidationError(f"{model_name} has empty state support")
|
|
753
|
+
initial = {state: state_mass[state] / total_state_mass for state in active_states}
|
|
754
|
+
|
|
755
|
+
flows: dict[tuple[Hashable, Any, Hashable], float] = defaultdict(float)
|
|
756
|
+
for pair, pair_mass in joint.items():
|
|
757
|
+
if pair_mass <= cutoff:
|
|
758
|
+
continue
|
|
759
|
+
for transition in bidir.graph.out_transitions(pair):
|
|
760
|
+
symbol = transition.data.get(ATTR_EMISSION)
|
|
761
|
+
prob = float(transition.data.get(ATTR_PROB, 0.0))
|
|
762
|
+
target_pair = transition.target
|
|
763
|
+
if symbol is None or prob <= cutoff or target_pair not in pair_channel:
|
|
764
|
+
continue
|
|
765
|
+
for state, state_prob in pair_channel[pair].items():
|
|
766
|
+
if state not in initial or state_prob <= cutoff:
|
|
767
|
+
continue
|
|
768
|
+
for target_state, target_prob in pair_channel[target_pair].items():
|
|
769
|
+
if target_state not in initial or target_prob <= cutoff:
|
|
770
|
+
continue
|
|
771
|
+
flows[(state, symbol, target_state)] += pair_mass * state_prob * prob * target_prob
|
|
772
|
+
|
|
773
|
+
graph = TransitionGraph()
|
|
774
|
+
for state in active_states:
|
|
775
|
+
graph.add_state(state)
|
|
776
|
+
|
|
777
|
+
by_source: dict[Hashable, dict[tuple[Hashable, Any], float]] = defaultdict(lambda: defaultdict(float))
|
|
778
|
+
for (source, symbol, target), flow in flows.items():
|
|
779
|
+
if flow > 0.0:
|
|
780
|
+
by_source[source][(target, symbol)] += flow
|
|
781
|
+
|
|
782
|
+
for source in active_states:
|
|
783
|
+
outgoing = by_source.get(source, {})
|
|
784
|
+
total = sum(outgoing.values())
|
|
785
|
+
if total <= cutoff:
|
|
786
|
+
raise StochasticValidationError(f"{model_name} state {source!r} has no outgoing mass")
|
|
787
|
+
added = False
|
|
788
|
+
for (target, symbol), flow in outgoing.items():
|
|
789
|
+
prob = flow / total
|
|
790
|
+
if prob <= cutoff:
|
|
791
|
+
continue
|
|
792
|
+
graph.add_transition(source, target, **{ATTR_PROB: float(prob), ATTR_EMISSION: symbol})
|
|
793
|
+
added = True
|
|
794
|
+
if not added:
|
|
795
|
+
target, symbol = max(outgoing, key=outgoing.__getitem__)
|
|
796
|
+
graph.add_transition(source, target, **{ATTR_PROB: 1.0, ATTR_EMISSION: symbol})
|
|
797
|
+
|
|
798
|
+
model = model_cls(
|
|
799
|
+
graph=graph,
|
|
800
|
+
initial_distribution=initial,
|
|
801
|
+
observation_alphabet=bidir.observation_alphabet,
|
|
802
|
+
pair_state_channel=pair_channel,
|
|
803
|
+
joint_pair_distribution=joint,
|
|
804
|
+
source_bidirectional=bidir,
|
|
805
|
+
entropy_rate=bidir.entropy_rate(),
|
|
806
|
+
**measure_kwargs,
|
|
807
|
+
)
|
|
808
|
+
model.validate()
|
|
809
|
+
return model
|
|
810
|
+
|
|
811
|
+
|
|
812
|
+
__all__ = [
|
|
813
|
+
"FunctionalGenerativeModel",
|
|
814
|
+
"GacsKornerGenerativeModel",
|
|
815
|
+
"MinimalGenerativeModel",
|
|
816
|
+
"WynerGenerativeModel",
|
|
817
|
+
"functional_generative_model",
|
|
818
|
+
"gacs_korner_generative_model",
|
|
819
|
+
"minimal_generative_model",
|
|
820
|
+
"wyner_generative_model",
|
|
821
|
+
]
|