Source code for pyphi.validate

# 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