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