Source code for knowledgespaces.estimation.constraints

"""Fixed and equality-constrained BLIM parameterizations.

The constrained M-step pools Bernoulli sufficient statistics within each
free equality class. Fixed classes do not contribute free parameters.
This is a direct maximization of the EM auxiliary function; see also
pks::blim's betafix/etafix and betaequal/etaequal interfaces.
"""

from __future__ import annotations

from collections.abc import Mapping, Sequence
from dataclasses import dataclass, field
from types import MappingProxyType

import numpy as np

from knowledgespaces.estimation.prediction import _state_vector


[docs] @dataclass(frozen=True) class BLIMConstraints: """Named restrictions shared by estimation, bootstrap and Jacobian. ``beta_fixed`` and ``eta_fixed`` map item labels to fixed probabilities in [0, 1). ``beta_equal`` and ``eta_equal`` contain groups of item labels constrained to share a probability. Overlapping groups merge transitively. Fixing any member fixes its entire equality class; conflicting fixed values in that class are rejected. ``pi_fixed``, when supplied, fixes the *complete* state distribution; its keys must match all states, including zero-probability ones. Unspecified item parameters and, by default, pi remain free. The beta + eta informative-item inequality is a separate restriction and is not imposed. Inputs are immutable snapshots and pickleable. """ beta_fixed: Mapping[str, float] = field(default_factory=dict) eta_fixed: Mapping[str, float] = field(default_factory=dict) beta_equal: Sequence[Sequence[str]] = () eta_equal: Sequence[Sequence[str]] = () pi_fixed: Mapping[frozenset[str], float] | None = None def __post_init__(self) -> None: for name in ("beta_fixed", "eta_fixed"): values = dict(getattr(self, name)) if any(not np.isfinite(v) or not 0 <= v < 1 for v in values.values()): raise ValueError(f"{name} values must be finite and in [0, 1).") object.__setattr__(self, name, MappingProxyType(values)) for name in ("beta_equal", "eta_equal"): groups = tuple(tuple(group) for group in getattr(self, name)) if any(not g or len(g) != len(set(g)) for g in groups): raise ValueError(f"{name} groups must be nonempty and contain unique items.") object.__setattr__(self, name, groups) if self.pi_fixed is not None: object.__setattr__(self, "pi_fixed", MappingProxyType(dict(self.pi_fixed))) def __reduce__(self) -> tuple: return ( type(self), ( dict(self.beta_fixed), dict(self.eta_fixed), self.beta_equal, self.eta_equal, None if self.pi_fixed is None else dict(self.pi_fixed), ), )
@dataclass class _ItemGroups: groups: list[np.ndarray] fixed: list[float | None] @property def n_free(self) -> int: return sum(value is None for value in self.fixed) def project(self, values: np.ndarray, *, clip_free: bool = True) -> None: for group, fixed in zip(self.groups, self.fixed, strict=True): mean = float(values[group].mean()) values[group] = ( (np.clip(mean, 1e-6, 1 - 1e-6) if clip_free else mean) if fixed is None else fixed ) def update( self, numerator: np.ndarray, denominator: np.ndarray, values: np.ndarray, *, eps_on_empty: bool, minimum_mass: float = 1e-10, ) -> None: for group, fixed in zip(self.groups, self.fixed, strict=True): if fixed is not None: values[group] = fixed else: total = float(denominator[group].sum()) if total > minimum_mass: values[group] = np.clip(numerator[group].sum() / total, 1e-6, 1 - 1e-6) elif eps_on_empty: values[group] = 1e-6 @dataclass class _CompiledConstraints: beta: _ItemGroups eta: _ItemGroups pi: np.ndarray | None def n_parameters(self, n_states: int) -> int: return self.beta.n_free + self.eta.n_free + (n_states - 1 if self.pi is None else 0) def update( self, W: np.ndarray, R: np.ndarray, S: np.ndarray, beta: np.ndarray, eta: np.ndarray, *, eps_on_empty: bool, observed: np.ndarray | None = None, minimum_mass: float = 1e-10, ) -> None: mastered = W @ S unmastered = W @ (1 - S) if observed is not None: mastered *= observed unmastered *= observed self.beta.update( ((1 - R) * mastered).sum(axis=0), mastered.sum(axis=0), beta, eps_on_empty=eps_on_empty, minimum_mass=minimum_mass, ) self.eta.update( (R * unmastered).sum(axis=0), unmastered.sum(axis=0), eta, eps_on_empty=eps_on_empty, minimum_mass=minimum_mass, ) def restrict_jacobian(self, jacobian: np.ndarray, n_items: int) -> np.ndarray: columns = [] for offset, spec in ((0, self.beta), (n_items, self.eta)): for group, fixed in zip(spec.groups, spec.fixed, strict=True): if fixed is None: columns.append(jacobian[:, offset + group].sum(axis=1)) if self.pi is None: columns.extend(jacobian[:, i] for i in range(2 * n_items, jacobian.shape[1])) return np.column_stack(columns) if columns else np.empty((len(jacobian), 0)) def _compile_item_groups( items: list[str], fixed: Mapping[str, float], equal: Sequence[Sequence[str]], ) -> _ItemGroups: known = set(items) if set(fixed) - known or any(set(group) - known for group in equal): raise ValueError("BLIM constraints reference unknown items.") parent = {q: q for q in items} def find(q: str) -> str: while parent[q] != q: parent[q] = parent[parent[q]] q = parent[q] return q for group in equal: for q in group[1:]: parent[find(q)] = find(group[0]) members: dict[str, list[int]] = {} for i, q in enumerate(items): members.setdefault(find(q), []).append(i) groups, values = [], [] for indices in members.values(): fixed_values = {float(fixed[items[i]]) for i in indices if items[i] in fixed} if len(fixed_values) > 1: raise ValueError("Conflicting fixed values in a BLIM equality group.") groups.append(np.array(indices)) values.append(next(iter(fixed_values)) if fixed_values else None) return _ItemGroups(groups, values) def _compile_constraints( constraints: BLIMConstraints, items: list[str], states: list[frozenset[str]], ) -> _CompiledConstraints: if not isinstance(constraints, BLIMConstraints): raise TypeError("constraints must be a BLIMConstraints instance.") return _CompiledConstraints( _compile_item_groups(items, constraints.beta_fixed, constraints.beta_equal), _compile_item_groups(items, constraints.eta_fixed, constraints.eta_equal), None if constraints.pi_fixed is None else _state_vector(constraints.pi_fixed, states), )