Source code for knowledgespaces.metrics.validation

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