# validate.py
"""Methods for validating user input."""
from __future__ import annotations
from collections.abc import Iterable
from collections.abc import Sequence
from typing import TYPE_CHECKING
from typing import Any
import numpy as np
from . import exceptions
from .conf import config
from .direction import Direction
if TYPE_CHECKING:
from pyphi.core.tpm.factored import FactoredTPM
# pylint: disable=redefined-outer-name
# TODO: move to `Direction`
[docs]
def directions(directions: Iterable[Direction], **kwargs: bool) -> bool:
"""Validate each direction in an iterable.
Parameters
----------
directions : Iterable[Direction]
Directions to validate.
**kwargs
Passed through to :func:`direction`.
Returns
-------
bool
``True`` if every element is a valid
:class:`~pyphi.direction.Direction`.
"""
return all(direction(d, **kwargs) for d in directions)
[docs]
def direction(direction: Direction, allow_bi: bool = False) -> bool:
"""Validate that the given direction is one of the allowed constants.
Parameters
----------
direction : Direction
Direction to validate.
allow_bi : bool
Whether the bidirectional arrow is also allowed.
Returns
-------
bool
``True`` if the direction is valid.
Raises
------
ValueError
If the direction is not one of the allowed constants.
"""
valid = set(Direction.both())
if allow_bi:
valid.add(Direction.BIDIRECTIONAL)
if direction not in valid:
raise ValueError(
f"`direction` must be one of `Direction.{valid}`; "
f"got {type(direction)} `{direction}`"
)
return True
[docs]
def connectivity_matrix(cm: np.ndarray) -> bool:
"""Validate the given connectivity matrix."""
# Special case for empty matrices.
if cm.size == 0:
return True
if cm.ndim != 2:
raise ValueError("Connectivity matrix must be 2-dimensional.")
if cm.shape[0] != cm.shape[1]:
raise ValueError("Connectivity matrix must be square.")
if not np.all(np.logical_or(cm == 1, cm == 0)):
raise ValueError("Connectivity matrix must contain only binary values.")
return True
[docs]
def connectivity(cm: np.ndarray, factored_tpm: FactoredTPM) -> bool:
"""Validate that a connectivity matrix is not under-specified.
An edge ``a -> b`` that the TPM implies (factor ``b`` depends on input
``a``) but ``cm`` omits would be silently marginalized out during node
construction and excluded from purview search, under-counting phi. Such
omissions are rejected. Over-specification (declaring an edge the TPM does
not use) is permitted — it only widens the search.
"""
inferred = factored_tpm.infer_cm()
cm_arr = np.asarray(cm)
missing = np.argwhere((inferred == 1) & (cm_arr == 0))
if missing.size:
edges = ", ".join(f"{int(a)} -> {int(b)}" for a, b in missing)
raise ValueError(
"Connectivity matrix is under-specified: the TPM implies "
f"edge(s) {edges} that the connectivity matrix omits. These "
"dependencies would be silently marginalized out, under-counting "
"phi. Add the edge(s) to the connectivity matrix, or set "
"validate_connectivity=False to allow a permissive matrix."
)
return True
[docs]
def node_labels(node_labels: Sequence[str], node_indices: Sequence[int]) -> None:
"""Validate that there is a label for each node."""
if len(node_labels) != len(node_indices):
raise ValueError(f"Labels {node_labels} must label every node {node_indices}.")
if len(node_labels) != len(set(node_labels)):
raise ValueError(f"Labels {node_labels} must be unique.")
[docs]
def substrate(n: object) -> bool:
"""Validate a :class:`~pyphi.substrate.Substrate`.
Checks the FactoredTPM and connectivity matrix.
"""
from pyphi.core.tpm.factored import FactoredTPM
factored = n.factored_tpm # type: ignore[attr-defined]
if not isinstance(factored, FactoredTPM):
raise ValueError("substrate.factored_tpm must be a FactoredTPM")
connectivity_matrix(n.cm) # type: ignore[attr-defined]
if n.cm.shape[0] != n.size: # type: ignore[attr-defined]
raise ValueError(
"Connectivity matrix must be NxN, where N is the "
"number of nodes in the substrate."
)
if config.infrastructure.validate_connectivity:
connectivity(n.cm, factored) # type: ignore[attr-defined]
return True
[docs]
def is_substrate(substrate: object) -> None:
"""Validate that the argument is a :class:`~pyphi.substrate.Substrate`."""
from . import Substrate
if not isinstance(substrate, Substrate):
raise ValueError(
"Input must be a Substrate (perhaps you passed a System instead?"
)
[docs]
def node_states(state: Sequence[int], alphabet_sizes: Sequence[int]) -> None:
"""Check that each state entry is a valid index into its node's alphabet.
Parameters
----------
state
Per-node state indices.
alphabet_sizes
Per-node alphabet sizes; ``state[i]`` must satisfy
``0 <= state[i] < alphabet_sizes[i]``.
"""
if len(state) != len(alphabet_sizes):
raise ValueError(
f"State length {len(state)} does not match alphabet_sizes length "
f"{len(alphabet_sizes)}."
)
for i, (s, k) in enumerate(zip(state, alphabet_sizes, strict=False)):
if not (0 <= s < k):
raise ValueError(
f"Invalid state: state[{i}]={s} is not in [0, {k}) for "
f"alphabet size {k}."
)
[docs]
def state_length(state: Sequence[int], size: int) -> bool:
"""Check that the state is the given size."""
if len(state) != size:
raise ValueError(
"Invalid state: there must be one entry per "
f"node in the substrate; this state has {len(state)} entries, but "
f"there are {size} nodes."
)
return True
[docs]
def transition_states(
substrate: Any,
before_state: Sequence[int],
after_state: Sequence[int],
) -> None:
"""Raise if the observed state pair is impossible under the dynamics.
Every unit's ``after_state`` must have nonzero probability given the
full ``before_state``: the Realization principle of Albantakis et al.
(2019) [1]_, Section 2.2, requires p(v_t | v_{t−1}) > 0 for a transition
to be defined.
Parameters
----------
substrate : Substrate
The substrate whose dynamics define transition probabilities.
before_state : tuple[int]
The state of the substrate at time t−1.
after_state : tuple[int]
The state of the substrate at time t.
Raises
------
pyphi.exceptions.TransitionUnreachableError
If ``p(after_state | before_state) = 0`` under the substrate's
factored TPM.
ValueError
If either state has the wrong length or is outside the alphabet.
References
----------
.. [1] Albantakis L, Marshall W, Hoel E, Tononi G. (2019).
What caused what? A quantitative account of actual causation using
dynamical causal networks. *Entropy*, 21 (5), 459.
https://doi.org/10.3390/e21050459
"""
factored = substrate.factored_tpm
state_length(before_state, substrate.size)
state_length(after_state, substrate.size)
node_states(before_state, factored.alphabet_sizes)
node_states(after_state, factored.alphabet_sizes)
for i in range(factored.n_nodes):
if factored._factor_at(i, before_state)[after_state[i]] <= 0.0:
raise exceptions.TransitionUnreachableError(
tuple(before_state), tuple(after_state)
)
[docs]
def state_reachable(system: object) -> None:
"""Raise :class:`~pyphi.exceptions.StateUnreachableForwardsError` if the
state is unreachable.
1. Substrate-level: the marginal probability
``P(state) = Σ_{s_t} ∏_i factor_i(s_t)[state[i]]`` must be positive
under the substrate's joint factored TPM.
2. Subsystem-level, under ``CONDITION_CURRENT_STATE`` only: the
subsystem's component of the state must be producible with the
background's past held at its current state. Some past state of the
*subsystem* must transition to the subsystem's ``proper_state`` with
nonzero probability.
Under ``CAUSAL_MARGINALIZATION``, check 1 is sufficient: the
background's past states are weighted by their probability given the
current state (Albantakis et al. 2023, Eq. 4), so if some past universe
state produces the current one, the cause TPM gives the system's state
positive probability from that past state. Holding the background fixed
on the cause side is the convention Albantakis et al. (2023) replaced
because it makes reachable states unreachable (S2 Text, "Background
Conditions").
"""
factored = system.substrate.factored_tpm # type: ignore[attr-defined]
state = system.state # type: ignore[attr-defined]
pr_joint = np.ones(factored.alphabet_sizes, dtype=np.float64)
for i in range(factored.n_nodes):
pr_joint = pr_joint * factored.factor(i)[..., state[i]]
if pr_joint.sum() <= 0.0:
raise exceptions.StateUnreachableForwardsError(system.state) # type: ignore[attr-defined]
# Subsystem-level: only when the background's past is held at its
# current state.
if (
system._resolved_background_conditioning() == "CONDITION_CURRENT_STATE" # type: ignore[attr-defined]
and not _proper_state_in_image_of_conditioned_tpm(system)
):
raise exceptions.StateUnreachableForwardsError(system.state) # type: ignore[attr-defined]
def _proper_state_in_image_of_conditioned_tpm(system: object) -> bool:
"""Whether the subsystem's ``proper_state`` is in the image of the
background-conditioned effect dynamics.
``proper_effect_marginal`` is a FactoredTPM with one factor per system
output unit (background fixed at the external state, background input
dims dropped). The state is in the image iff some system-input
configuration assigns positive joint probability to ``proper_state`` —
i.e. every system factor gives positive probability to its component
of ``proper_state`` for that input. Works for any per-unit alphabet
size.
"""
proper = system.proper_effect_marginal # type: ignore[attr-defined]
proper_state = system.proper_state # type: ignore[attr-defined]
joint = np.ones(proper.alphabet_sizes, dtype=np.float64)
for slot in range(proper.n_nodes):
joint = joint * proper.factor(slot)[..., proper_state[slot]]
return bool(np.any(joint > 0.0))
[docs]
def system_partition(partition: object, node_indices: Sequence[int]) -> None:
"""Check that the partition covers only the given nodes."""
if set(partition.indices) != set(node_indices): # type: ignore[attr-defined]
raise ValueError(
f"{partition} nodes are not equal to system nodes {node_indices}"
)
[docs]
def system(s: object) -> bool:
"""Validate a :class:`~pyphi.system.System`.
Checks its state and partition.
"""
node_states(s.state, s.substrate.factored_tpm.alphabet_sizes) # type: ignore[attr-defined]
system_partition(s.partition, s.partition_indices) # type: ignore[attr-defined]
if config.infrastructure.validate_system_states:
state_reachable(s)
return True
[docs]
def relata(relata: Iterable[object] | None) -> None:
"""Validate a set of relata."""
if not relata:
raise ValueError("relata cannot be empty")
[docs]
def non_overlapping(complexes: Iterable[Any]) -> bool:
"""Validate that complexes have pairwise-disjoint units (exclusion).
The exclusion postulate requires that no unit belongs to more than one
complex.
Parameters
----------
complexes : Iterable
Objects exposing ``node_indices``.
Returns
-------
bool
``True`` if the complexes are pairwise node-disjoint.
Raises
------
ValueError
If any two of ``complexes`` share a unit.
"""
seen: set[int] = set()
for c in complexes:
units = set(c.node_indices or ())
overlap = units & seen
if overlap:
raise ValueError(
f"Exclusion violated: unit(s) {sorted(overlap)} belong to more "
f"than one complex."
)
seen.update(units)
return True