"""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),
)