"""The triggered TPM: the system's fixed-lag response to each stimulus."""
from __future__ import annotations
import itertools
from dataclasses import dataclass
import numpy as np
import pandas as pd
from pyphi import convert
from pyphi import utils
from pyphi.labels import NodeLabels
[docs]
@dataclass(frozen=True)
class TriggeredTPM:
"""Pr(Sₜ = s | ∂S_{t−τ} = x), one distribution over system states per
stimulus.
Attributes
----------
array : numpy.ndarray
A multidimensional array with one binary axis per unit, ordered
``(sensory axes..., system axes...)``; ``array[x + s]`` is
Pr(S = s | ∂S = x). Marginalizing a unit subset is a uniform axis sum.
sensory_indices : tuple of int
Substrate indices of the sensory-interface units, in axis order.
system_indices : tuple of int
Substrate indices of the system units, in axis order.
node_labels : NodeLabels
Labels for the substrate units, for the labeled ``to_pandas`` view.
"""
array: np.ndarray
sensory_indices: tuple[int, ...]
system_indices: tuple[int, ...]
node_labels: NodeLabels
[docs]
def row(self, stimulus: tuple[int, ...]) -> np.ndarray:
"""The system-state distribution for one stimulus."""
return self.array[tuple(stimulus)]
[docs]
def argmax_state(self, stimulus: tuple[int, ...]) -> tuple[int, ...]:
"""The most-probable system state for a stimulus (the triggered state).
Ties resolve to the first maximum in little-endian state order.
"""
flat = int(np.argmax(self.row(stimulus).ravel(order="F")))
return convert.le_index2state(flat, len(self.system_indices))
def _marginalize_system(self, distribution, mechanism, state) -> float:
"""Return Pr(mechanism = state) from a distribution over the system axes.
Sums out the system units not in ``mechanism``. Requires ``mechanism``
to be a subset of ``system_indices`` (without duplicates) and ``state``
to match its length; the (mechanism, state) pairs may be given in any
order.
"""
mechanism = tuple(mechanism)
if len(set(mechanism)) != len(mechanism):
raise ValueError(f"duplicate units in mechanism {mechanism}")
if not set(mechanism) <= set(self.system_indices):
raise ValueError(
f"mechanism {mechanism} is not a subset of system_indices "
f"{self.system_indices}"
)
if len(state) != len(mechanism):
raise ValueError(f"state {state} length != mechanism {mechanism} length")
# Canonicalize: sort the (mechanism, state) pairs together so the
# axis bookkeeping below can assume increasing mechanism order.
pairs = sorted(zip(mechanism, state, strict=True))
mechanism = tuple(m for m, _ in pairs)
state = tuple(s for _, s in pairs)
keep = [self.system_indices.index(m) for m in mechanism]
sum_axes = tuple(a for a in range(len(self.system_indices)) if a not in keep)
reduced = distribution.sum(axis=sum_axes) if sum_axes else distribution
# `mechanism` is sorted above and `system_indices` is validated sorted
# at construction, so `keep` is increasing and the remaining axes are
# already in mechanism order.
return float(reduced[tuple(state)])
[docs]
def conditional_probability(self, mechanism, state, stimulus) -> float:
"""Pr(mechanism = state | ∂S = stimulus)."""
return self._marginalize_system(self.row(stimulus), mechanism, state)
[docs]
def marginal_probability(self, mechanism, state) -> float:
"""Pr(mechanism = state), the uniform-prior marginal over stimuli."""
marginal = self.array.mean(axis=tuple(range(len(self.sensory_indices))))
return self._marginalize_system(marginal, mechanism, state)
[docs]
def to_pandas(self) -> pd.DataFrame:
"""Labeled view: rows = stimulus states, columns = system states,
values = Pr(s | x)."""
from pyphi.models.pandas import state_multiindex
index = state_multiindex(self.node_labels, self.sensory_indices)
columns = state_multiindex(self.node_labels, self.system_indices)
data = [[self.array[tuple(x) + tuple(s)] for s in columns] for x in index]
return pd.DataFrame(data, index=index, columns=columns)
def _full_state(sensory_indices, system_indices, x, s_sys, n):
full = [0] * n
for i, xi in zip(sensory_indices, x, strict=True):
full[i] = xi
for i, si in zip(system_indices, s_sys, strict=True):
full[i] = si
return tuple(full)
def _validate_binary_substrate(substrate) -> None:
"""Raise if the substrate has any non-binary unit.
The triggered-TPM construction operates on the binary state-by-node
representation; only binary substrates are currently supported.
"""
sizes = substrate.factored_tpm.alphabet_sizes
if any(size != 2 for size in sizes):
raise ValueError(
f"only binary substrates are currently supported; got alphabet sizes {sizes}"
)
def _validate_sorted_indices(name: str, indices) -> None:
"""Raise unless ``indices`` is strictly increasing (sorted, no duplicates).
Triggered-TPM axes and stimulus/state tuples are positional relative
to these index tuples, so only the sorted form is unambiguous.
"""
if not all(a < b for a, b in itertools.pairwise(indices)):
raise ValueError(
f"{name} must be strictly increasing (sorted, without "
f"duplicates); got {tuple(indices)}"
)
def _system_step_tpm(sbn_full, sensory_indices, system_indices, n, *, clamp_to):
"""A one-step state-by-node TPM over the system, with the sensory interface
either clamped to a state (``clamp_to=x``) or marginalized
(``clamp_to=None``)."""
system = list(system_indices)
shape_s = (2,) * len(system_indices)
step = np.zeros((*shape_s, len(system_indices)))
for s_sys in utils.all_states(len(system_indices)):
if clamp_to is not None:
full = _full_state(sensory_indices, system_indices, clamp_to, s_sys, n)
step[s_sys] = sbn_full[full][system]
else:
acc = np.zeros(len(system_indices))
for x in utils.all_states(len(sensory_indices)):
full = _full_state(sensory_indices, system_indices, x, s_sys, n)
acc += sbn_full[full][system]
step[s_sys] = acc / (2 ** len(sensory_indices))
return step
def _lagged_sbs(step_sbn, t):
sbs = convert.state_by_node2state_by_state(step_sbn)
if t == 0:
return np.eye(sbs.shape[0])
return np.linalg.matrix_power(sbs, t)
[docs]
def build_triggered_tpm(
substrate, sensory_indices, system_indices, *, tau, tau_clamp
) -> TriggeredTPM:
"""Construct the triggered TPM by clamp-then-noise evolution.
Clamp the sensory interface to the stimulus for ``tau_clamp`` steps, then
marginalize it for the remaining ``tau - tau_clamp`` steps; compose and
average over the initial system state. Only binary substrates are
currently supported.
"""
_validate_binary_substrate(substrate)
_validate_sorted_indices("sensory_indices", sensory_indices)
_validate_sorted_indices("system_indices", system_indices)
n = len(substrate.node_indices)
sbn_full = np.asarray(substrate.tpm.to_array())[..., 1] # binary ON-prob slice
noised = _lagged_sbs(
_system_step_tpm(sbn_full, sensory_indices, system_indices, n, clamp_to=None),
tau - tau_clamp,
)
rows = []
for x in utils.all_states(len(sensory_indices)):
clamped = _lagged_sbs(
_system_step_tpm(sbn_full, sensory_indices, system_indices, n, clamp_to=x),
tau_clamp,
)
composed = clamped @ noised
rows.append(composed.mean(axis=0)) # marginalize initial system state
flat = np.array(rows) # (n_stimuli, n_system_states), little-endian flat order
n_sensory, n_system = len(sensory_indices), len(system_indices)
array = flat.reshape((2,) * (n_sensory + n_system))
# The flat orders are little-endian (first unit varies fastest) but the
# C-order reshape unpacks last-axis-fastest, leaving each axis group in
# reversed unit order; transpose each group back to unit order.
sensory_axes = tuple(reversed(range(n_sensory)))
system_axes = tuple(n_sensory + a for a in reversed(range(n_system)))
return TriggeredTPM(
array=array.transpose(sensory_axes + system_axes),
sensory_indices=tuple(sensory_indices),
system_indices=tuple(system_indices),
node_labels=substrate.node_labels,
)