"""BLIM cell diagnostics over the complete multinomial response table."""
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
if TYPE_CHECKING:
from knowledgespaces.estimation.blim_em import BLIMEstimate, ResponseMatrix
[docs]
@dataclass(frozen=True)
class BLIMResiduals:
"""Complete response table and signed cell residuals.
Patterns use lexicographic binary order with columns in ``items``
(the supplied data's order). Duplicate rows are pooled and all absent
patterns receive observed count zero. Arrays are read-only snapshots.
Pearson residuals are (O-E)/sqrt(E); deviance residuals are
sign(O-E)*sqrt(2*(O*log(O/E)-O+E)), taking 0*log(0/E)=0.
A cell with O=E=0 contributes zero to both diagnostics; O>0 with
model probability exactly zero contributes +inf. Log probabilities
distinguish impossible cells from floating-point underflow.
The sums of squares give G2 and X2 on the full table, independent of
degrees-of-freedom conventions. They do not establish a chi-square
reference distribution, identify causes of misfit, or account for
testing many individual cells. Residuals can also assess held-out data.
"""
items: tuple[str, ...]
patterns: np.ndarray
observed: np.ndarray
expected: np.ndarray
log_probabilities: np.ndarray
pearson: np.ndarray
deviance: np.ndarray
@property
def G2(self) -> float:
"""Sum of squared deviance residuals."""
return float(np.dot(self.deviance, self.deviance))
@property
def X2(self) -> float:
"""Sum of squared Pearson residuals."""
return float(np.dot(self.pearson, self.pearson))
[docs]
def blim_residuals(
estimate: BLIMEstimate,
data: ResponseMatrix,
*,
chunk_size: int = 1024,
max_patterns: int = 1_048_576,
max_memory_bytes: int = 512_000_000,
) -> BLIMResiduals:
"""Evaluate Pearson and deviance residuals, including unobserved cells.
No refit is performed; parameters align by labels, allowing held-out
data and reordered columns. Counts may be fractional analysis weights;
such weights do not automatically admit multinomial inference.
Enumeration is exponential in the number of items. Chunking limits
prediction temporaries but the returned complete table remains in
memory. The memory preflight estimates allocations, not process RSS.
"""
if set(data.items) != set(estimate.items):
raise ValueError("data.items must match the estimate's domain.")
chunk_size = _positive_integer(chunk_size, "chunk_size")
max_memory_bytes = _positive_integer(max_memory_bytes, "max_memory_bytes")
q, m = data.n_items, len(estimate.states)
count = _pattern_count(q, max_patterns)
batch = min(count, chunk_size)
table_bytes = 8 * (count * (q + 16) + data.n_patterns * (q + 2))
prediction_bytes = 8 * (6 * batch * m + 3 * batch * q + 4 * m * q)
if table_bytes + prediction_bytes > max_memory_bytes:
raise MemoryError("Residuals exceed max_memory_bytes; reduce chunk_size or domain size.")
patterns = _binary_patterns(0, count, q)
codes = data.patterns.astype(np.int64) @ (2 ** np.arange(q - 1, -1, -1))
observed = np.bincount(codes, weights=data.effective_counts, minlength=count)
log_probabilities = np.empty(count)
for start in range(0, count, chunk_size):
prediction = estimate.predict(
patterns[start : start + chunk_size],
items=data.items,
max_memory_bytes=max_memory_bytes - table_bytes,
)
log_probabilities[start : start + chunk_size] = prediction.log_probabilities
log_expected = np.log(data.n_respondents) + log_probabilities
expected = np.exp(log_expected)
# Empty cells have negative residuals, including a zero limit at E=0.
pearson = -np.exp(0.5 * log_expected)
deviance = -np.sqrt(2 * expected)
positive = observed > 0
o, e, le = observed[positive], expected[positive], log_expected[positive]
difference = o - e
contribution = o * (np.log(o) - le) - difference
# For nearly fitted cells use log1p to reduce cancellation in O log(O/E).
near = (e > 0) & (np.abs(difference) < 0.1 * e)
ratio = difference[near] / e[near]
contribution[near] = e[near] * ((1 + ratio) * np.log1p(ratio) - ratio)
deviance[positive] = np.sign(difference) * np.sqrt(2 * np.maximum(contribution, 0))
with np.errstate(divide="ignore", over="ignore"):
pearson[positive] = np.sign(difference) * np.exp(np.log(np.abs(difference)) - 0.5 * le)
arrays = (patterns, observed, expected, log_probabilities, pearson, deviance)
for array in arrays:
array.flags.writeable = False
return BLIMResiduals(tuple(data.items), *arrays)