"""BLIM prediction on specified response patterns, including exact zeros."""
from __future__ import annotations
from collections.abc import Mapping
from dataclasses import dataclass
from typing import TYPE_CHECKING, Literal
import numpy as np
from scipy.special import logsumexp
from knowledgespaces._patterns import _positive_integer
from knowledgespaces.estimation.discrepancy import (
MDOptions,
_discrepancy_weights,
_validate_options,
)
from knowledgespaces.structures.knowledge_structure import KnowledgeStructure
if TYPE_CHECKING:
from knowledgespaces.estimation.constraints import BLIMConstraints
ItemParameter = float | Mapping[str, float] | np.ndarray
StateParameter = Mapping[frozenset[str], float] | np.ndarray | None
def _item_vector(value: ItemParameter, items: list[str], name: str) -> np.ndarray:
if isinstance(value, Mapping):
if set(value) != set(items):
raise ValueError(f"{name} keys must match the item domain exactly.")
vector = np.array([value[q] for q in items], dtype=float)
else:
vector = np.asarray(value, dtype=float)
if vector.ndim == 0:
vector = np.full(len(items), float(vector))
if vector.shape != (len(items),):
raise ValueError(f"{name} must contain one probability per item.")
if not np.isfinite(vector).all() or np.any((vector < 0) | (vector > 1)):
raise ValueError(f"{name} values must be finite and in [0, 1].")
return vector.copy()
def _state_vector(value: StateParameter, states: list[frozenset[str]]) -> np.ndarray:
if value is None:
return np.full(len(states), 1 / len(states))
if isinstance(value, Mapping):
if set(value) != set(states):
raise ValueError("pi keys must match the knowledge states exactly.")
vector = np.array([value[s] for s in states], dtype=float)
else:
vector = np.asarray(value, dtype=float)
if (
vector.shape != (len(states),)
or not np.isfinite(vector).all()
or np.any(vector < 0)
or not np.isclose(vector.sum(), 1, rtol=0, atol=1e-8)
):
raise ValueError("pi must be a finite, nonnegative distribution summing to 1.")
return vector / vector.sum()
def _state_matrix(items: list[str], states: list[frozenset[str]]) -> np.ndarray:
return np.array([[q in state for q in items] for state in states], dtype=float)
def _log_conditional(patterns: np.ndarray, correct: np.ndarray) -> np.ndarray:
"""Log P(R|K); avoid the undefined product 0 * log(0)."""
with np.errstate(divide="ignore"):
log_yes = np.log(correct)
log_no = np.log1p(-correct)
if np.all((correct > 0) & (correct < 1)):
return patterns @ log_yes.T + (1 - patterns) @ log_no.T
result = np.zeros((len(patterns), len(correct)))
for q in range(patterns.shape[1]):
result += np.where(patterns[:, q, None] == 1, log_yes[None, :, q], log_no[None, :, q])
return result
[docs]
@dataclass(frozen=True)
class BLIMPrediction:
"""Predictions in input row order and canonical state order.
``probabilities`` and ``log_probabilities`` describe P(R).
``posteriors[r, k]`` is P(K_k|R_r). A model-impossible pattern has
probability zero, log probability -inf, and an all-NaN posterior:
conditioning on a zero-probability event is undefined.
These Bayesian quantities are unchanged by the assignment ``method``.
``state_probabilities`` gives the selected ML/MD/MDML assignments;
only ML is the ordinary BLIM posterior. MD ignores priors and error
magnitudes, apart from declared fixed-zero error exclusions. MDML
renormalizes the joint probability on the included states.
Arrays are read-only snapshots. Use ``classify`` explicitly to select
one state per row; undefined assignments produce None, not the empty set.
"""
items: tuple[str, ...]
states: tuple[frozenset[str], ...]
probabilities: np.ndarray
log_probabilities: np.ndarray
posteriors: np.ndarray
method: Literal["ML", "MD", "MDML"] = "ML"
discrepancy_assignments: np.ndarray | None = None
@property
def state_probabilities(self) -> np.ndarray:
"""Row-normalized assignments for the requested prediction method."""
return (
self.posteriors
if self.discrepancy_assignments is None
else self.discrepancy_assignments
)
[docs]
def classify(
self,
*,
ties: Literal["min", "max", "random"] = "min",
seed: int | np.random.Generator | None = None,
) -> tuple[frozenset[str] | None, ...]:
"""Select modal states, with exact computed-probability ties.
``min``/``max`` select smallest/largest cardinality among modal
states, then the first in ``states`` order (canonical for the
prediction functions). ``random`` samples uniformly among all
modal states, using a local NumPy generator; no R/NumPy random
stream identity is implied. Singleton maxima do not consume draws.
Each input row receives one selection, irrespective of frequency.
An all-NaN assignment row returns None.
"""
if ties not in ("min", "max", "random"):
raise ValueError("ties must be 'min', 'max', or 'random'.")
if seed is not None and ties != "random":
raise ValueError("seed applies only to random tie-breaking.")
rng = np.random.default_rng(seed) if ties == "random" else None
result: list[frozenset[str] | None] = []
sizes = np.array([len(state) for state in self.states])
for row in self.state_probabilities:
if not np.isfinite(row).all():
result.append(None)
continue
candidates = np.flatnonzero(row == row.max())
if len(candidates) > 1:
if rng is not None:
chosen = int(rng.choice(candidates))
else:
target = sizes[candidates].min() if ties == "min" else sizes[candidates].max()
chosen = int(candidates[sizes[candidates] == target][0])
else:
chosen = int(candidates[0])
result.append(self.states[chosen])
return tuple(result)
def _constraint_feasibility(
constraints: BLIMConstraints | None,
items: list[str],
states: list[frozenset[str]],
patterns: np.ndarray,
membership: np.ndarray,
beta: np.ndarray,
eta: np.ndarray,
prior: np.ndarray,
) -> np.ndarray | None:
"""Validate declared restrictions and exclude fixed-zero error events."""
if constraints is None:
return None
# Local import avoids the parameter-validation import cycle.
from knowledgespaces.estimation.constraints import _compile_constraints
compiled = _compile_constraints(constraints, items, states)
feasible = np.ones((len(patterns), len(states)), dtype=bool)
for name, values, groups in (("beta", beta, compiled.beta), ("eta", eta, compiled.eta)):
for group, fixed in zip(groups.groups, groups.fixed, strict=True):
target = values[group[0]] if fixed is None else fixed
if not np.all(values[group] == target):
raise ValueError(f"{name} values do not satisfy the declared constraints.")
if fixed == 0:
for q in group:
error = (
(patterns[:, q, None] == 0) & (membership[None, :, q] == 1)
if name == "beta"
else (patterns[:, q, None] == 1) & (membership[None, :, q] == 0)
)
feasible &= ~error
if compiled.pi is not None and not np.allclose(prior, compiled.pi, rtol=0, atol=1e-14):
raise ValueError("pi values do not satisfy the declared constraints.")
return feasible
[docs]
def predict_blim(
structure: KnowledgeStructure,
patterns: np.ndarray,
*,
items: list[str] | None = None,
beta: ItemParameter = 0.1,
eta: ItemParameter = 0.1,
pi: StateParameter = None,
method: Literal["ML", "MD", "MDML"] = "ML",
discrepancy: MDOptions | None = None,
inclusion: np.ndarray | None = None,
constraints: BLIMConstraints | None = None,
max_memory_bytes: int = 512_000_000,
) -> BLIMPrediction:
"""Predict BLIM response probabilities and state posteriors.
``patterns`` is a binary row-by-item matrix. ``items`` defaults to
the sorted domain and determines the column order, including that
of array-valued beta/eta. Mappings align by labels. State arrays
use canonical order (size, then lexicographic). Responses are
conditionally independent given a fixed knowledge state.
Accepts the closed probability box [0, 1], including deterministic
limits and uninformative items. The informative-item inequality
beta + eta < 1 is not imposed by this probability calculation.
No clipping of zero probabilities or imputation of missing values.
The memory preflight is an allocation estimate, not an OS-level cap.
``method`` selects state assignments (default ML, independently of how
parameters were fitted). MD normalizes discrepancy weights; MDML
normalizes the joint BLIM probabilities restricted to included states.
``discrepancy`` supplies the excess Hamming radius, or hyperbolic weights
for MD only. Alternatively, ``inclusion`` is an explicit binary
row-by-canonical-state mask (pks i.RK), mutually exclusive with discrepancy.
An empty selection has undefined (all-NaN) assignments. Bayesian P(R)
and ``posteriors`` always retain their original full-model meaning.
Declared ``constraints`` must match supplied parameters. Fixed-zero
beta/eta (including equality groups) exclude incompatible state-response
pairs from computed discrepancy assignments, as in pks. Undeclared zero
values do not change the MD distance rule. A supplied inclusion mask is
used as-is, independently of the distance rule and those exclusions;
MDML still respects zero joint probabilities. Zero prior mass does not
exclude a state from MD. No data-dependent fitting is performed here.
"""
if method not in ("ML", "MD", "MDML"):
raise ValueError("method must be 'ML', 'MD', or 'MDML'.")
spec = _validate_options(discrepancy, method)
if inclusion is not None and (method == "ML" or discrepancy is not None):
raise ValueError("inclusion requires MD/MDML and cannot be combined with discrepancy.")
items = sorted(structure.domain) if items is None else list(items)
if len(items) != len(set(items)) or set(items) != structure.domain:
raise ValueError("items must list the structure's domain exactly once.")
patterns = np.asarray(patterns)
if patterns.ndim != 2 or patterns.shape[1] != len(items):
raise ValueError("patterns must be a 2D matrix with one column per item.")
if not np.isin(patterns, (0, 1)).all():
raise ValueError("patterns must contain only binary responses (0 or 1).")
max_memory_bytes = _positive_integer(max_memory_bytes, "max_memory_bytes")
states = list(structure)
n, m, q = len(patterns), len(states), len(items)
matrices = 6 if method == "ML" and constraints is None else 12
if 8 * (matrices * n * m + 3 * n * q + 4 * m * q) > max_memory_bytes:
raise MemoryError("predict_blim exceeds max_memory_bytes; predict fewer rows per call.")
b = _item_vector(beta, items, "beta")
e = _item_vector(eta, items, "eta")
prior = _state_vector(pi, states)
membership = _state_matrix(items, states)
feasible = _constraint_feasibility(
constraints, items, states, patterns, membership, b, e, prior
)
correct = membership * (1 - b) + (1 - membership) * e
joint = _log_conditional(patterns.astype(float), correct)
with np.errstate(divide="ignore"):
joint += np.log(prior)
log_probability = logsumexp(joint, axis=1)
posterior = np.full(joint.shape, np.nan)
possible = np.isfinite(log_probability)
posterior[possible] = np.exp(joint[possible] - log_probability[possible, None])
probability = np.exp(log_probability)
assignments = None
if method != "ML":
if inclusion is None:
weights = _discrepancy_weights(
patterns.astype(float), membership, np.zeros(n), spec, feasible
)
if feasible is not None:
# Estimation ignores zero-frequency impossible rows. Prediction
# instead keeps their assignments undefined, without inventing a state.
weights[~feasible.any(axis=1)] = 0
else:
weights = np.asarray(inclusion)
if weights.shape != (n, m) or not np.isin(weights, (0, 1)).all():
raise ValueError("inclusion must be a binary pattern-by-state matrix.")
weights = weights.astype(float)
assignments = np.full((n, m), np.nan)
if method == "MD":
total = weights.sum(axis=1)
assigned = total > 0
assignments[assigned] = weights[assigned] / total[assigned, None]
else:
restricted = np.where(weights > 0, joint, -np.inf)
normalizer = logsumexp(restricted, axis=1)
assigned = np.isfinite(normalizer)
assignments[assigned] = np.exp(restricted[assigned] - normalizer[assigned, None])
assignments.flags.writeable = False
for array in (probability, log_probability, posterior):
array.flags.writeable = False
return BLIMPrediction(
tuple(items), tuple(states), probability, log_probability, posterior, method, assignments
)