Source code for knowledgespaces.estimation.prediction

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