"""Empirical validation of structures and prerequisite relations.
Direct implementations of the definitions documented by kst::kvalidate
(Schrepp, 1999; Schrepp, Held & Albert, 1999). These are descriptive
data–structure indices, distinct from model-based BLIM reliability.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING
import numpy as np
from knowledgespaces._patterns import _binary_patterns, _pattern_count
from knowledgespaces.estimation.prediction import _positive_integer, _state_matrix
from knowledgespaces.structures.knowledge_structure import KnowledgeStructure
from knowledgespaces.structures.relations import SurmiseRelation
if TYPE_CHECKING:
from knowledgespaces.estimation.blim_em import ResponseMatrix
def _nearest_distances(patterns: np.ndarray, membership: np.ndarray) -> np.ndarray:
rows = patterns.astype(float)
distances = rows.sum(axis=1)[:, None] + membership.sum(axis=1)
distances -= 2 * (rows @ membership.T)
return distances.min(axis=1).astype(np.int64)
[docs]
@dataclass(frozen=True)
class StructureValidation:
"""Empirical distance summary, with read-only arrays.
``pattern_distances`` follows the input rows, including zero-weight
rows. ``distance_frequencies[d]`` is their total weight at distance d,
including d=0. ``di`` is their weighted mean distance in items.
``uniform_distance`` is the mean over the complete response powerset;
``da = di / uniform_distance``. Both it and its frequency table are
None when ``compute_da=False``. DA is NaN for a powerset structure
(zero denominator), and can exceed one. Smaller DI/DA means closer
responses; neither is a significance test or a probability of validity.
"""
items: tuple[str, ...]
pattern_distances: np.ndarray
distance_frequencies: np.ndarray
di: float
da: float | None
uniform_distance: float | None
uniform_distance_frequencies: np.ndarray | None
[docs]
def validate_structure(
structure: KnowledgeStructure,
data: ResponseMatrix,
*,
compute_da: bool = True,
chunk_size: int = 1024,
max_patterns: int = 1_048_576,
max_memory_bytes: int = 512_000_000,
) -> StructureValidation:
"""Nearest Hamming distances, DI, and optionally exact DA.
The supplied family need not be union closed and is not enlarged.
Binary responses align by ``data.items``. Integer counts and fractional
nonnegative analysis weights are supported. DA compares DI with the
uniform mean distance over all 2**|Q| responses, not a fitted BLIM.
Chunking bounds temporary storage, not exponential enumeration time.
``max_patterns`` guards DA enumeration only; ``max_memory_bytes`` is
a conservative allocation estimate, not a process memory guarantee.
"""
if set(data.items) != structure.domain:
raise ValueError("data.items must match the structure's domain.")
chunk_size = _positive_integer(chunk_size, "chunk_size")
max_memory_bytes = _positive_integer(max_memory_bytes, "max_memory_bytes")
q, m, n = data.n_items, len(structure), data.n_patterns
count = _pattern_count(q, max_patterns) if compute_da else 0
batch = min(chunk_size, max(n, count))
if 8 * (4 * batch * m + 4 * batch * q + 3 * m * q + n + 4 * (q + 1)) > max_memory_bytes:
raise MemoryError("Validation exceeds max_memory_bytes; reduce chunk_size.")
membership = _state_matrix(data.items, list(structure))
distances = np.empty(n, dtype=np.int64)
for start in range(0, n, chunk_size):
distances[start : start + chunk_size] = _nearest_distances(
data.patterns[start : start + chunk_size], membership
)
frequencies = np.bincount(distances, weights=data.effective_counts, minlength=q + 1)
di = float(np.dot(np.arange(q + 1), frequencies / data.n_respondents))
uniform_frequencies = None
uniform_distance = da = None
if compute_da:
uniform_frequencies = np.zeros(q + 1, dtype=np.int64)
for start in range(0, count, chunk_size):
patterns = _binary_patterns(start, min(start + chunk_size, count), q)
d = _nearest_distances(patterns, membership)
uniform_frequencies += np.bincount(d, minlength=q + 1)
uniform_distance = float(np.dot(np.arange(q + 1), uniform_frequencies / count))
da = di / uniform_distance if uniform_distance else float("nan")
uniform_frequencies.flags.writeable = False
distances.flags.writeable = frequencies.flags.writeable = False
return StructureValidation(
tuple(data.items), distances, frequencies, di, da, uniform_distance, uniform_frequencies
)
[docs]
@dataclass(frozen=True)
class RelationValidation:
"""Descriptive oriented-pair counts and coefficients.
``pairs`` lists every nonreflexive pair (a, b): b implies a. Each
respondent contributes once per pair, so counts may exceed sample size.
A=1,B=0 is concordant; A=0,B=1 is discordant; ties contribute neither.
``gamma = (n_concordant - n_discordant)/(n_concordant + n_discordant)``;
it is NaN without untied responses. ``vc`` divides discordances by
total response weight times the number of pairs; it is NaN without
nonreflexive pairs. ``solution_rates`` are fractions in input item
order, not percentages. The array is read-only.
"""
items: tuple[str, ...]
pairs: tuple[tuple[str, str], ...]
n_concordant: float
n_discordant: float
gamma: float
vc: float
solution_rates: np.ndarray
[docs]
def validate_relation(
relation: SurmiseRelation | KnowledgeStructure, data: ResponseMatrix
) -> RelationValidation:
"""Compute KST gamma, VC and item solution rates.
Relation generators are transitively closed before counting; equivalent
items contribute both directed pairs. A structure is replaced by its
implied surmise relation. For non-quasi-ordinal structures this tests
only the pairwise prerequisites, not all restrictions of the family;
use :func:`validate_structure` to assess the family itself.
"""
if isinstance(relation, KnowledgeStructure):
relation = relation.surmise_relation()
if set(data.items) != relation.items:
raise ValueError("data.items must match the relation's domain.")
pairs = tuple(sorted(relation.transitive_closure().relations))
columns = {item: i for i, item in enumerate(data.items)}
weights = data.effective_counts
concordant = discordant = 0.0
for a, b in pairs:
x, y = data.patterns[:, columns[a]], data.patterns[:, columns[b]]
concordant += float(weights[(x == 1) & (y == 0)].sum())
discordant += float(weights[(x == 0) & (y == 1)].sum())
untied = concordant + discordant
rates = (weights / data.n_respondents) @ data.patterns
rates.flags.writeable = False
return RelationValidation(
tuple(data.items),
pairs,
concordant,
discordant,
(concordant - discordant) / untied if untied else float("nan"),
discordant / data.n_respondents / len(pairs) if pairs else float("nan"),
rates,
)