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.
Files changed (150) hide show
  1. sofic/__init__.py +185 -0
  2. sofic/automata/__init__.py +207 -0
  3. sofic/automata/_config_simulation.py +40 -0
  4. sofic/automata/active.py +611 -0
  5. sofic/automata/alergia.py +222 -0
  6. sofic/automata/algorithms.py +376 -0
  7. sofic/automata/atomaton.py +58 -0
  8. sofic/automata/base.py +161 -0
  9. sofic/automata/buchi.py +23 -0
  10. sofic/automata/buchi_simulation.py +67 -0
  11. sofic/automata/canonical_dual.py +18 -0
  12. sofic/automata/canonical_extraction.py +122 -0
  13. sofic/automata/dfa.py +85 -0
  14. sofic/automata/dfasat.py +195 -0
  15. sofic/automata/edsm.py +219 -0
  16. sofic/automata/enumeration.py +44 -0
  17. sofic/automata/icdfa.py +421 -0
  18. sofic/automata/idfa.py +363 -0
  19. sofic/automata/languages/__init__.py +39 -0
  20. sofic/automata/languages/_quotient_utils.py +64 -0
  21. sofic/automata/languages/atoms.py +31 -0
  22. sofic/automata/languages/automaton_ops.py +243 -0
  23. sofic/automata/languages/base.py +67 -0
  24. sofic/automata/languages/operations.py +78 -0
  25. sofic/automata/languages/quotients.py +66 -0
  26. sofic/automata/languages/residuals.py +25 -0
  27. sofic/automata/learning.py +79 -0
  28. sofic/automata/nfa.py +39 -0
  29. sofic/automata/nwa.py +343 -0
  30. sofic/automata/nwa_simulation.py +56 -0
  31. sofic/automata/observation.py +40 -0
  32. sofic/automata/papni.py +301 -0
  33. sofic/automata/regex.py +128 -0
  34. sofic/automata/rfsa.py +35 -0
  35. sofic/automata/rpni.py +193 -0
  36. sofic/automata/subsequential.py +201 -0
  37. sofic/automata/transducer_operations.py +350 -0
  38. sofic/automata/transducer_simulation.py +150 -0
  39. sofic/automata/transducers.py +365 -0
  40. sofic/automata/unifilar.py +107 -0
  41. sofic/automata/vpa.py +1373 -0
  42. sofic/automata/vpa_simulation.py +53 -0
  43. sofic/base.py +153 -0
  44. sofic/core.py +47 -0
  45. sofic/examples/__init__.py +86 -0
  46. sofic/examples/epsilon_machines.py +1089 -0
  47. sofic/examples/processes.py +1491 -0
  48. sofic/examples/shifts.py +144 -0
  49. sofic/exceptions.py +33 -0
  50. sofic/generators/__init__.py +115 -0
  51. sofic/generators/_word_measures.py +94 -0
  52. sofic/generators/alternative_complexity.py +104 -0
  53. sofic/generators/base.py +327 -0
  54. sofic/generators/bidirectional_construction.py +717 -0
  55. sofic/generators/bidirectional_epsilon_machine.py +689 -0
  56. sofic/generators/block_convergence.py +668 -0
  57. sofic/generators/block_entropy.py +578 -0
  58. sofic/generators/channel_measures.py +75 -0
  59. sofic/generators/conversions.py +182 -0
  60. sofic/generators/directional_flow.py +245 -0
  61. sofic/generators/edge_emissions.py +36 -0
  62. sofic/generators/edge_machine.py +178 -0
  63. sofic/generators/epsilon_construction.py +193 -0
  64. sofic/generators/epsilon_inference.py +703 -0
  65. sofic/generators/epsilon_machine.py +557 -0
  66. sofic/generators/epsilon_transducer.py +168 -0
  67. sofic/generators/epsilon_transducer_construction.py +185 -0
  68. sofic/generators/epsilon_transducer_inference.py +499 -0
  69. sofic/generators/hmm_inference.py +719 -0
  70. sofic/generators/information_diagram.py +428 -0
  71. sofic/generators/lumping.py +447 -0
  72. sofic/generators/markov.py +100 -0
  73. sofic/generators/mealy.py +156 -0
  74. sofic/generators/measures.py +257 -0
  75. sofic/generators/minimal_generative_model.py +821 -0
  76. sofic/generators/mixed_state.py +250 -0
  77. sofic/generators/mixed_state_construction.py +163 -0
  78. sofic/generators/moore.py +75 -0
  79. sofic/generators/nmachine.py +78 -0
  80. sofic/generators/nmachine_construction.py +70 -0
  81. sofic/generators/pfa.py +100 -0
  82. sofic/generators/prob.py +291 -0
  83. sofic/generators/process_equivalence.py +207 -0
  84. sofic/generators/quasi_inference.py +74 -0
  85. sofic/generators/quasi_realization.py +97 -0
  86. sofic/generators/reversal.py +66 -0
  87. sofic/generators/stack_hmm.py +426 -0
  88. sofic/generators/stack_inference.py +509 -0
  89. sofic/generators/stationary.py +134 -0
  90. sofic/generators/stochastic.py +65 -0
  91. sofic/generators/synchronization.py +407 -0
  92. sofic/generators/topological_epsilon_enumeration.py +349 -0
  93. sofic/generators/words.py +226 -0
  94. sofic/graph.py +135 -0
  95. sofic/indexing.py +31 -0
  96. sofic/inference/__init__.py +45 -0
  97. sofic/inference/bayesian/__init__.py +68 -0
  98. sofic/inference/bayesian/comparison.py +199 -0
  99. sofic/inference/bayesian/counts.py +219 -0
  100. sofic/inference/bayesian/diversity.py +254 -0
  101. sofic/inference/bayesian/epsilon.py +270 -0
  102. sofic/inference/bayesian/hdp_hmm.py +340 -0
  103. sofic/inference/bayesian/markov.py +294 -0
  104. sofic/inference/bayesian/pymc_backend.py +71 -0
  105. sofic/inference/bayesian/stack_hmm.py +215 -0
  106. sofic/inference/model_selection.py +365 -0
  107. sofic/inference/spectral.py +564 -0
  108. sofic/operations.py +16 -0
  109. sofic/properties.py +339 -0
  110. sofic/serialization.py +450 -0
  111. sofic/shifts/__init__.py +48 -0
  112. sofic/shifts/algorithms.py +84 -0
  113. sofic/shifts/base.py +49 -0
  114. sofic/shifts/cover_construction.py +76 -0
  115. sofic/shifts/covers.py +47 -0
  116. sofic/shifts/dyck_algorithms.py +100 -0
  117. sofic/shifts/dyck_enumeration.py +275 -0
  118. sofic/shifts/markov_dyck.py +172 -0
  119. sofic/shifts/parry_construction.py +82 -0
  120. sofic/shifts/sft.py +104 -0
  121. sofic/shifts/sft_construction.py +52 -0
  122. sofic/shifts/sliding_block_code.py +156 -0
  123. sofic/shifts/sofic.py +111 -0
  124. sofic/shifts/sofic_dyck.py +110 -0
  125. sofic/shifts/sofic_relation.py +64 -0
  126. sofic/shifts/textile.py +104 -0
  127. sofic/shifts/tmc.py +46 -0
  128. sofic/shifts/tmc_construction.py +58 -0
  129. sofic/shifts/topological_anatomy.py +150 -0
  130. sofic/states.py +27 -0
  131. sofic/testing/__init__.py +8 -0
  132. sofic/testing/strategies.py +154 -0
  133. sofic/viz/__init__.py +16 -0
  134. sofic/viz/_context.py +345 -0
  135. sofic/viz/_edge.py +216 -0
  136. sofic/viz/_format.py +89 -0
  137. sofic/viz/_labels.py +34 -0
  138. sofic/viz/_names.py +17 -0
  139. sofic/viz/_rational.py +20 -0
  140. sofic/viz/_tikz_compile.py +177 -0
  141. sofic/viz/_tikz_format.py +122 -0
  142. sofic/viz/_tikz_layout.py +218 -0
  143. sofic/viz/assets/vaucanson.tikz +71 -0
  144. sofic/viz/graphviz.py +158 -0
  145. sofic/viz/idiagram.py +350 -0
  146. sofic/viz/tikz.py +381 -0
  147. sofic-0.1.0.dist-info/METADATA +444 -0
  148. sofic-0.1.0.dist-info/RECORD +150 -0
  149. sofic-0.1.0.dist-info/WHEEL +4 -0
  150. 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
+ ]