Source code for knowledgespaces.metrics.reliability

"""Full-test BLIM reliability from de Chiusole et al. (2024, 2025).

2024: doi:10.3758/s13428-024-02468-3, Equations 10-11.
2025: doi:10.1111/bmsp.70013, Section 3, Equations 6-9.
Direct finite-sum implementations, with response enumeration in chunks.
"""

from __future__ import annotations

from dataclasses import dataclass

import numpy as np
from scipy.special import xlogy

from knowledgespaces.estimation.prediction import (
    ItemParameter,
    StateParameter,
    _item_vector,
    _log_conditional,
    _positive_integer,
    _state_matrix,
    _state_vector,
)
from knowledgespaces.structures.knowledge_structure import KnowledgeStructure


[docs] @dataclass(frozen=True) class BLIMReliability: """Population quantities for a fixed BLIM and administration of all items. ``rp_reliability`` = I(K;R)/H(R); ``ks_reliability`` = I(K;R)/H(K). Entropies and mutual information use bits. Zero entropy in a denominator produces NaN. These coefficients concern the model, not a particular respondent or an adaptive stopping rule. ``accuracy`` is P(MAP state = true state), with uniform random choice among tied posterior modes. ``chance_accuracy`` is max(pi), and ``adjusted_accuracy`` is (accuracy - chance)/(1 - chance), the 2025 kappa-prime index; it is NaN if chance is one. ``discrepancy_distribution[d]`` is ``P(|K xor MAP| = d)``, including d=0; ``expected_discrepancy`` is its unconditional mean in items. ``state_accuracy[k]`` is conditional accuracy given state k; it is NaN for states with zero prior mass. Arrays are read-only snapshots. """ rp_reliability: float ks_reliability: float mutual_information: float entropy_responses: float entropy_states: float conditional_entropy_responses: float accuracy: float chance_accuracy: float adjusted_accuracy: float expected_discrepancy: float discrepancy_distribution: np.ndarray states: tuple[frozenset[str], ...] state_accuracy: np.ndarray n_response_patterns: int tie_tolerance: float
[docs] def blim_reliability( structure: KnowledgeStructure, *, beta: ItemParameter = 0.1, eta: ItemParameter = 0.1, pi: StateParameter = None, chunk_size: int = 1024, max_patterns: int = 1_048_576, max_memory_bytes: int = 512_000_000, tie_tolerance: float = 1e-12, ) -> BLIMReliability: """Compute the 2024 entropy and 2025 MAP reliability indices. Parameters align with the sorted item domain and canonical state order; mappings are recommended for fitted parameters. Closed-box probabilities [0, 1] permit deterministic and independence limits. No sample is simulated: all 2**|Q| patterns are summed in bounded-size chunks. This reduces memory use, not the exponential enumeration time. ``tie_tolerance`` groups log joint probabilities within this absolute distance of their maximum as tied modes. Set zero for equality of the computed values. Tied modes receive equal classification probability, as specified in the 2025 paper; no random seed is needed. ``max_patterns`` limits work before enumeration; ``max_memory_bytes`` bounds a conservative allocation estimate, not actual process memory. These are full-test, fixed-model quantities. Plugging in a fit does not account for parameter uncertainty or establish validity for adaptive assessment or unobserved responses. """ chunk_size = _positive_integer(chunk_size, "chunk_size") max_patterns = _positive_integer(max_patterns, "max_patterns") max_memory_bytes = _positive_integer(max_memory_bytes, "max_memory_bytes") if not np.isfinite(tie_tolerance) or tie_tolerance < 0: raise ValueError("tie_tolerance must be finite and nonnegative.") items, states = sorted(structure.domain), list(structure) q, m = len(items), len(states) n_patterns = 2**q if n_patterns > max_patterns: raise ValueError(f"Reliability needs {n_patterns} patterns, exceeding max_patterns.") if q >= 64: raise ValueError("The response enumerator supports at most 63 items.") batch = min(chunk_size, n_patterns) if 8 * (8 * batch * m + 3 * batch * q + 6 * m * q + 4 * m) > max_memory_bytes: raise MemoryError("Reliability exceeds max_memory_bytes; reduce chunk_size.") b = _item_vector(beta, items, "beta") e = _item_vector(eta, items, "eta") prior = _state_vector(pi, states) membership = _state_matrix(items, states) correct = membership * (1 - b) + (1 - membership) * e conditional_entropy = float( prior @ (-xlogy(correct, correct) - xlogy(1 - correct, 1 - correct)).sum(axis=1) ) / np.log(2) state_entropy = float(-xlogy(prior, prior).sum()) / np.log(2) with np.errstate(divide="ignore"): log_prior = np.log(prior) entropy = 0.0 distribution = np.zeros(q + 1) state_accuracy = np.zeros(m) for start in range(0, n_patterns, batch): numbers = np.arange(start, min(start + batch, n_patterns), dtype=np.uint64) patterns = ((numbers[:, None] >> np.arange(q, dtype=np.uint64)[::-1]) & 1).astype(float) log_conditional = _log_conditional(patterns, correct) log_joint = log_conditional + log_prior joint = np.exp(log_joint) marginal = joint.sum(axis=1) entropy -= float(xlogy(marginal, marginal).sum()) / np.log(2) possible = np.isfinite(log_joint.max(axis=1)) # Impossible responses have zero joint mass, and no assigned mode. modes = np.zeros(joint.shape) maxima = log_joint[possible].max(axis=1, keepdims=True) modes[possible] = maxima - log_joint[possible] <= tie_tolerance totals = modes.sum(axis=1, keepdims=True) np.divide(modes, totals, out=modes, where=totals > 0) state_accuracy += (np.exp(log_conditional) * modes).sum(axis=0) # Sum by estimated state; no |K| by |K| distance/confusion matrix. for predicted in np.flatnonzero(modes.any(axis=0)): masses = modes[:, predicted] @ joint distances = np.count_nonzero(membership != membership[predicted], axis=1) distribution += np.bincount(distances, weights=masses, minlength=q + 1) distribution /= distribution.sum() # remove accumulated floating-point drift information = float(np.clip(entropy - conditional_entropy, 0, min(entropy, state_entropy))) chance = float(prior.max()) accuracy = float(distribution[0]) state_accuracy[prior == 0] = np.nan state_accuracy.flags.writeable = False distribution.flags.writeable = False return BLIMReliability( rp_reliability=information / entropy if entropy > 0 else float("nan"), ks_reliability=information / state_entropy if state_entropy > 0 else float("nan"), mutual_information=information, entropy_responses=entropy, entropy_states=state_entropy, conditional_entropy_responses=float(conditional_entropy), accuracy=accuracy, chance_accuracy=chance, adjusted_accuracy=(accuracy - chance) / (1 - chance) if chance < 1 else float("nan"), expected_discrepancy=float(distribution @ np.arange(q + 1)), discrepancy_distribution=distribution, states=tuple(states), state_accuracy=state_accuracy, n_response_patterns=n_patterns, tie_tolerance=tie_tolerance, )