diff --git a/README.md b/README.md index 6874766..7e83ad0 100644 --- a/README.md +++ b/README.md @@ -167,7 +167,9 @@ $ python main_optimization.py --task odd-one-out \ ### NOTES: -1. Note that triplet data is expected to be in the format `N x 3`, where `N` = number of triplets (e.g., 100k) and 3 refers to the three objects in a triplet, where `col_0` = anchor, `col_1` = positive, `col_2` = odd-one-out/negative. Triplet data must be split into `train` and `test` splits, and named `train_90.txt` or `train_90.npy` and `test_10.txt` or `test_10.npy` respectively. +1. Note that triplet data is expected to be in the format `N x 3`, where `N` = number of triplets (e.g., 100k) and 3 refers to the three objects in a triplet, where `col_0` = anchor, `col_1` = positive, `col_2` = odd-one-out/negative. Triplet data must be split into `train` and `test` splits, and named `train_90.txt` or `train_90.npy` and `test_10.txt` or `test_10.npy` respectively. Every file that is loaded is checked for repeated triplets, and a `DuplicateTripletsWarning` is issued if there are more of them than the sampling design can explain. Some repeats are always expected: if `N` trials are drawn from `M` objects, the birthday paradox predicts a definite number of collisions, and a clean sample of 200k triplets over the 1854 THINGS objects contains around 38 repeated rows. So the check reports the observed count, the count expected under a null in which the `N` trials are drawn uniformly from the `K = C(M,3)` possible triplets, and the probability of seeing at least that many; it warns only when that probability falls below `duplicate_stats.P_VALUE_THRESHOLD` (0.2). Two triplets count as the same trial when they use the same three objects, whatever order the columns are in. `M` is taken from `--n_objects` where that flag exists and inferred from the data otherwise. The warning never changes the data that is loaded, and it can be silenced with `warnings.filterwarnings("ignore", category=utils.DuplicateTripletsWarning)`. + + What this does and does not tell you: the test flags repeats that uniform sampling cannot account for. It cannot distinguish *deliberately* repeated trials — the same triplet shown to several participants, which `partition_triplets.py` preserves on purpose — from a file that was accidentally concatenated twice, because neither is random. Both will trip it. What it buys you is silence in the regime where repeats are genuinely expected, which is where an unconditional duplicate check is pure noise. Lower `P_VALUE_THRESHOLD` if you are working with deliberately repeated designs. 2. Every `--steps` epochs (i.e., `if (epoch + 1) % steps == 0`) a `model_epoch.tar` (including model and optimizer `state_dicts`) and a `results_epoch.json` (including train and validation cross-entropy errors) file are saved to disk. In addition, after convergence of VICE, a `pruned_params.npz` (compressed binary file) with keys `pruned_loc` and `pruned_scale`, including pruned VICE parameters, is saved to disk. Latent dimensions of the pruned parameter matrices are sorted according to their overall importance. See output folder structure below for where to find these files.
diff --git a/duplicate_stats.py b/duplicate_stats.py new file mode 100644 index 0000000..537a2d5 --- /dev/null +++ b/duplicate_stats.py @@ -0,0 +1,338 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +"""Detect triplets that occur more often in the data than chance would predict. + +Repeated trials are not by themselves a problem: if ``N`` triplets are drawn at +random from ``M`` objects, the birthday paradox guarantees a predictable number +of repeats. A clean uniform sample of 200,000 triplets over the 1,854 THINGS +objects contains around 38 repeated rows. Warning about those is noise. + +What is worth a warning is a number of repeats that the sampling design cannot +explain -- a file that was concatenated with itself, say. So this module tests +the observed count against an explicit null instead of thresholding it. + +Null model +---------- +The ``N`` trials are drawn independently and uniformly from the +``K = C(M, 3)`` possible triplets. Trial identity is the *unordered* set +``{i, j, k}``: the same three objects shown twice is the same trial, whichever +column each object landed in. (``partition_triplets.py`` uses the same +convention.) The null is deliberately not taken over ordered rows, because the +odd-one-out choice is the signal VICE models, not a uniform draw over three +options -- two participants shown the same triplet agree far more often than +one time in three. + +Test statistic +-------------- +``W``, the number of *pairs* of trials that carry the same triplet. Writing +``I_ab`` for the indicator that trials ``a`` and ``b`` collide, the ``I_ab`` are +pairwise uncorrelated: for distinct ``a, b, c``, +``E[I_ab I_ac] = 1/K^2 = E[I_ab] E[I_ac]``. So ``W`` matches +``Binomial(C(N, 2), 1/K)`` in both moments exactly, + + E[W] = C(N, 2) / K Var[W] = C(N, 2) (1/K) (1 - 1/K) + +which is what makes a binomial reference distribution exactly right here rather +than merely convenient. The p-value is the one-sided upper tail. Only an excess +of repeats is flagged: a deficit means sampling without replacement, which is +what ``main_tripletize.py`` does on purpose. +""" + +import math +import warnings +from dataclasses import dataclass +from typing import Optional, Tuple + +import torch +from scipy.special import gammaln +from scipy.stats import binom + +Tensor = torch.Tensor + +# Warn only when the observed repeats fall in this upper tail under the null. +# Monkeypatch this rather than threading a keyword through load_data. +# +# Worth knowing where this bites: whenever fewer than -log(1 - 0.2) = 0.223 +# collisions are expected, P(W >= 1) is already below the threshold and the +# test reduces exactly to "warn on any repeat at all". Sparse designs therefore +# keep the unconditional behaviour; only the dense regime, where repeats are +# genuinely predicted, goes quiet. +P_VALUE_THRESHOLD = 0.2 + +_LOG10 = math.log(10.0) + + +class DuplicateTripletsWarning(UserWarning): + """Raised when the triplet data repeats triplets more often than chance.""" + + +@dataclass(frozen=True) +class DuplicateReport: + """Observed and expected repeat counts for one triplet file. + + Invariant: ``p_value == 0.0`` implies ``log10_p_value < -300``. The p-value + underflows float64 long before the evidence stops being interpretable, so + read ``log10_p_value`` rather than taking a log of ``p_value``. + """ + + file_name: str + n_rows: int # N + n_objects: Optional[int] # M + n_cells: int # K = C(M, 3) + n_colliding_pairs: int # W, the test statistic + n_duplicated_rows: int # rows sharing their triplet with another row + n_duplicated_triplets: int # distinct triplets that repeat + n_duplicated_rows_exact: int # identical *rows*, reported not tested + n_degenerate_rows: int # rows whose three indices are not distinct + expected_colliding_pairs: float + expected_duplicated_rows: float + expected_duplicated_triplets: float + p_value: float + log10_p_value: float + testable: bool # False when the null is undefined for this input + + @property + def is_significant(self) -> bool: + """Whether this report warrants a warning.""" + if self.n_duplicated_rows == 0: + return False + # An undefined null must not silence the check: fall back to warning + # about any repeat at all, without the statistical claim. + if not self.testable: + return True + return self.p_value < P_VALUE_THRESHOLD + + def summary(self) -> str: + """The warning body.""" + head = ( + f"\n...Found {self.n_duplicated_rows} of {self.n_rows} rows in " + f"'{self.file_name}' that repeat one of {self.n_duplicated_triplets} " + "distinct triplets.\n" + ) + if self.testable: + head += ( + f"...Chance alone predicts {self.expected_duplicated_rows:.1f} such " + f"rows for M = {self.n_objects} objects and N = {self.n_rows} trials " + f"(K = {self.n_cells:,} possible triplets): " + f"{self.n_colliding_pairs} colliding pairs observed vs " + f"{self.expected_colliding_pairs:.1f} expected, " + f"{_format_pvalue(self.p_value, self.log10_p_value)}.\n" + ) + else: + head += ( + "...The expected number of repeats could not be computed for " + f"M = {self.n_objects}, so this is a plain count with no test " + "behind it.\n" + ) + if self.n_degenerate_rows: + head += ( + f"...{self.n_degenerate_rows} rows do not contain three distinct " + "objects; those rows cannot arise under the null and are counted " + "but not modelled.\n" + ) + head += ( + "...Repeated trials are expected if the same triplet was shown to " + "several participants; a rate far above chance usually means a file " + "was concatenated twice. VICE trains on the data exactly as loaded.\n" + ) + return head + + +def _format_pvalue(p_value: float, log10_p_value: float) -> str: + if math.isnan(p_value): + return "p could not be computed" + if p_value >= 1e-300: + return f"p = {p_value:.2g} under a binomial null" + return f"p < 1e-300 (log10 p = {log10_p_value:.1f}) under a binomial null" + + +def n_cells(n_objects: int) -> int: + """Number of distinct triplets that can be drawn from n_objects items.""" + if n_objects is None or n_objects < 3: + return 0 + return math.comb(n_objects, 3) + + +def count_duplicates(triplets: Tensor) -> Tuple[int, int]: + """Count rows that are *exactly* equal to another row. + + Column order is significant here, so ``[0, 1, 2]`` and ``[2, 1, 0]`` are + different rows. Returns the number of rows belonging to a duplicated group + and the number of distinct rows those collapse into. Reported for the + user's benefit; the statistical test works on unordered triplets instead. + """ + _, counts = torch.unique(triplets, dim=0, return_counts=True) + repeated = counts[counts > 1] + return int(repeated.sum().item()), int(repeated.shape[0]) + + +def collision_counts(triplets: Tensor) -> Tuple[int, int, int]: + """Count repeats of the unordered triplet {i, j, k}. + + Returns ``(n_colliding_pairs, n_duplicated_rows, n_duplicated_triplets)``. + Sorting each row before the unique pass is the whole difference between + row identity and trial identity. + """ + if triplets.shape[0] == 0: + return 0, 0, 0 + sorted_rows = torch.sort(triplets.long(), dim=1).values + _, counts = torch.unique(sorted_rows, dim=0, return_counts=True) + repeated = counts[counts > 1] + n_colliding_pairs = int((counts * (counts - 1) // 2).sum().item()) + return n_colliding_pairs, int(repeated.sum().item()), int(repeated.shape[0]) + + +def count_degenerate_rows(triplets: Tensor) -> int: + """Count rows whose three indices are not distinct, e.g. [5, 5, 7].""" + if triplets.shape[0] == 0: + return 0 + a, b, c = triplets[:, 0], triplets[:, 1], triplets[:, 2] + return int(((a == b) | (a == c) | (b == c)).sum().item()) + + +def expected_duplicates(n_rows: int, n_cells: int) -> Tuple[float, float, float]: + """Expected repeat counts under the null, in numerically stable form. + + Returns ``(colliding_pairs, duplicated_rows, duplicated_triplets)``. + + The closed forms below are algebraically equal to the textbook occupancy + expressions but are written with ``log1p``/``expm1`` throughout. This is not + cosmetic: the direct form of the third quantity, + ``K (1 - q**N - N p q**(N-1))``, subtracts numbers that agree to 16 digits + when ``N**2 << K``. At ``M = 20000, N = 50`` it returns 5.28e-05 against a + true value of 9.19e-10, and at ``M = 5000, N = 100`` it returns a negative + count. + + Some cancellation survives in the third quantity even so: at the sparse + extreme it holds about six significant digits rather than the full sixteen. + That is far more than a warning message needs, and the p-value -- the only + number the warning actually decides on -- does not depend on it. + """ + N, K = n_rows, n_cells + if N < 2 or K <= 0: + return 0.0, 0.0, 0.0 + q = math.log1p(-1.0 / K) + expected_pairs = math.comb(N, 2) / K + expected_rows = -N * math.expm1((N - 1) * q) + expected_triplets = -K * math.expm1((N - 1) * q + math.log1p((N - 1) / K)) + return expected_pairs, expected_rows, expected_triplets + + +def collision_pvalue( + n_colliding_pairs: int, n_rows: int, n_cells: int +) -> Tuple[float, float]: + """One-sided upper-tail probability of the observed collision count. + + Returns ``(p_value, log10_p_value)``. ``p_value`` underflows to 0.0 past + roughly 1e-308; ``log10_p_value`` stays finite well beyond that, because the + interesting cases (a file concatenated with itself) sit thousands of orders + of magnitude into the tail and "p = 0.0" tells the user nothing. + """ + N, K, w = n_rows, n_cells, n_colliding_pairs + # No collisions is never surprising, whatever the null looks like -- and + # this is also the N < 2 case, where there are no pairs to collide. + if w <= 0: + return 1.0, 0.0 + if N < 2 or K <= 0: + return float("nan"), float("nan") + n_pairs = math.comb(N, 2) + p_value = float(binom.sf(w - 1, n_pairs, 1.0 / K)) + if p_value > 0.0: + return p_value, math.log10(p_value) + # binom.sf has underflowed, and binom.logsf/poisson.logsf return -inf here + # rather than the log we need. Fall back to the Poisson tail in log space: + # the leading term plus a geometric-ratio correction. Accurate to about + # four decimals in log10 wherever the exact value is still computable. + lam = n_pairs / K + if w > lam > 0.0: + log_tail = (-lam + w * math.log(lam) - float(gammaln(w + 1))) - math.log1p( + -lam / w + ) + return 0.0, log_tail / _LOG10 + return 0.0, float("-inf") + + +def analyse( + triplets: Tensor, n_objects: Optional[int] = None, file_name: str = "" +) -> DuplicateReport: + """Build a DuplicateReport for one triplet array.""" + testable = True + + if triplets.dim() != 2 or triplets.shape[1] != 3: + # Not triplet-shaped; count nothing rather than silently comparing + # scalars, but do not raise from what is only a data check. + return DuplicateReport( + file_name=file_name, + n_rows=int(triplets.shape[0]) if triplets.dim() else 0, + n_objects=n_objects, + n_cells=0, + n_colliding_pairs=0, + n_duplicated_rows=0, + n_duplicated_triplets=0, + n_duplicated_rows_exact=0, + n_degenerate_rows=0, + expected_colliding_pairs=0.0, + expected_duplicated_rows=0.0, + expected_duplicated_triplets=0.0, + p_value=float("nan"), + log10_p_value=float("nan"), + testable=False, + ) + + n_rows = int(triplets.shape[0]) + if n_rows and int(triplets.min().item()) < 0: + # Negative indices are not object labels, so C(M, 3) is not the sample + # space. Still report the repeats, just without the test. + testable = False + + K = n_cells(n_objects) if n_objects is not None else 0 + if K <= 0: + testable = False + + pairs, dup_rows, dup_triplets = collision_counts(triplets) + dup_rows_exact, _ = count_duplicates(triplets) if n_rows else (0, 0) + degenerate = count_degenerate_rows(triplets) + + if testable: + exp_pairs, exp_rows, exp_triplets = expected_duplicates(n_rows, K) + p_value, log10_p = collision_pvalue(pairs, n_rows, K) + else: + exp_pairs = exp_rows = exp_triplets = float("nan") + p_value = log10_p = float("nan") + + return DuplicateReport( + file_name=file_name, + n_rows=n_rows, + n_objects=n_objects, + n_cells=K, + n_colliding_pairs=pairs, + n_duplicated_rows=dup_rows, + n_duplicated_triplets=dup_triplets, + n_duplicated_rows_exact=dup_rows_exact, + n_degenerate_rows=degenerate, + expected_colliding_pairs=exp_pairs, + expected_duplicated_rows=exp_rows, + expected_duplicated_triplets=exp_triplets, + p_value=p_value, + log10_p_value=log10_p, + testable=testable, + ) + + +def warn_on_duplicates( + triplets: Tensor, file_name: str, n_objects: Optional[int] = None +) -> DuplicateReport: + """Warn if the data repeats triplets more often than chance would predict. + + Never alters or de-duplicates the data, and never raises: a data check must + not be able to stop a training run. + """ + report = analyse(triplets, n_objects=n_objects, file_name=file_name) + if report.is_significant: + warnings.warn( + report.summary(), + DuplicateTripletsWarning, + stacklevel=3, + ) + return report diff --git a/main_inference.py b/main_inference.py index 6964d0d..b9c6c98 100644 --- a/main_inference.py +++ b/main_inference.py @@ -214,10 +214,10 @@ def inference( ) ) test_triplets = utils.load_data( - device=device, triplets_dir=triplets_dir, inference=True + device=device, triplets_dir=triplets_dir, inference=True, n_objects=n_objects ) _, val_triplets = utils.load_data( - device=device, triplets_dir=triplets_dir, inference=False + device=device, triplets_dir=triplets_dir, inference=False, n_objects=n_objects ) val_triplets = TripletData( triplets=val_triplets, diff --git a/main_robustness_eval.py b/main_robustness_eval.py index 80ceb28..77f5e3b 100644 --- a/main_robustness_eval.py +++ b/main_robustness_eval.py @@ -230,7 +230,7 @@ def evaluate_models( ) model_paths = get_model_paths(in_path) _, val_triplets = utils.load_data( - device=device, triplets_dir=triplets_dir, inference=False + device=device, triplets_dir=triplets_dir, inference=False, n_objects=n_objects ) val_triplets = TripletData( triplets=val_triplets, n_objects=n_objects, diff --git a/tests/test_duplicate_stats.py b/tests/test_duplicate_stats.py new file mode 100644 index 0000000..4868b15 --- /dev/null +++ b/tests/test_duplicate_stats.py @@ -0,0 +1,267 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +"""Tests for the duplicate-triplet statistics. + +Everything here is deterministic and exact: the reference values come from +fractions.Fraction, not from Monte Carlo, so the suite stays fast and cannot +flake in CI. +""" + +import math +import unittest +from collections import Counter +from fractions import Fraction + +import duplicate_stats as ds +import torch +from scipy.stats import binom + + +def exact_expected_duplicated_triplets(n_rows: int, n_cells: int) -> float: + """K (1 - q^N - N p q^(N-1)) in exact rational arithmetic.""" + p = Fraction(1, n_cells) + q = 1 - p + return float(n_cells * (1 - q ** n_rows - n_rows * p * q ** (n_rows - 1))) + + +def naive_expected_duplicated_triplets(n_rows: int, n_cells: int) -> float: + """The textbook form, written the obvious way in floating point.""" + N, K = n_rows, n_cells + q = math.log1p(-1.0 / K) + return K * (1 - math.exp(N * q) - (N / K) * math.exp((N - 1) * q)) + + +def exact_upper_tail(w: int, n: int, n_cells: int) -> float: + """Exact binomial upper tail, for small n.""" + p = Fraction(1, n_cells) + return float(sum(math.comb(n, k) * p ** k * (1 - p) ** (n - k) for k in range(w, n + 1))) + + +class ExpectedCountsTestCase(unittest.TestCase): + def test_expected_pairs_is_exact(self) -> None: + for n_rows, n_objects in [(85, 20), (1500, 1854), (200000, 1854)]: + K = ds.n_cells(n_objects) + pairs, _, _ = ds.expected_duplicates(n_rows, K) + self.assertEqual(pairs, math.comb(n_rows, 2) / K) + + def test_stable_form_matches_exact_rational(self) -> None: + # Six significant digits at the sparse extreme, where some cancellation + # survives even the stable form; far better elsewhere. + for n_rows, n_objects in [(100, 1854), (1000, 1854), (50, 20000), (100, 5000)]: + K = ds.n_cells(n_objects) + _, _, stable = ds.expected_duplicates(n_rows, K) + exact = exact_expected_duplicated_triplets(n_rows, K) + self.assertLess(abs(stable / exact - 1.0), 1e-5) + + def test_naive_form_is_why_this_is_written_the_hard_way(self) -> None: + """Lock in the stable form: the obvious algebra is catastrophically wrong. + + If someone 'simplifies' expected_duplicates back to the textbook + expression, this test is what tells them why it was not written that + way. At M=5000, N=100 the naive form returns a negative expected count; + at M=20000, N=50 it is wrong by five orders of magnitude. + """ + K = ds.n_cells(5000) + self.assertLess(naive_expected_duplicated_triplets(100, K), 0.0) + _, _, stable = ds.expected_duplicates(100, K) + self.assertGreater(stable, 0.0) + + K = ds.n_cells(20000) + exact = exact_expected_duplicated_triplets(50, K) + naive = naive_expected_duplicated_triplets(50, K) + _, _, stable = ds.expected_duplicates(50, K) + self.assertLess(abs(stable / exact - 1.0), 1e-5) + self.assertGreater(naive / exact, 1e4) + + def test_stable_form_agrees_with_naive_where_naive_works(self) -> None: + """The stable form is not merely different; it agrees when both are valid.""" + K = 100 + _, _, stable = ds.expected_duplicates(10, K) + self.assertAlmostEqual(stable, naive_expected_duplicated_triplets(10, K), places=9) + self.assertAlmostEqual(stable, exact_expected_duplicated_triplets(10, K), places=9) + + def test_expected_rows_is_about_twice_expected_pairs_when_sparse(self) -> None: + # each collision contributes two rows when triple collisions are negligible + K = ds.n_cells(1854) + pairs, rows, _ = ds.expected_duplicates(1500, K) + self.assertAlmostEqual(rows / (2 * pairs), 1.0, places=5) + + def test_degenerate_inputs_return_zero(self) -> None: + self.assertEqual(ds.expected_duplicates(1, 100), (0.0, 0.0, 0.0)) + self.assertEqual(ds.expected_duplicates(0, 100), (0.0, 0.0, 0.0)) + self.assertEqual(ds.expected_duplicates(100, 0), (0.0, 0.0, 0.0)) + + +class PValueTestCase(unittest.TestCase): + def test_matches_exact_rational_tail(self) -> None: + n_rows, n_objects, w = 10, 6, 3 + K = ds.n_cells(n_objects) + p, _ = ds.collision_pvalue(w, n_rows, K) + self.assertAlmostEqual(p, exact_upper_tail(w, math.comb(n_rows, 2), K), places=12) + + def test_zero_collisions_is_certain(self) -> None: + p, log10p = ds.collision_pvalue(0, 1000, ds.n_cells(1854)) + self.assertEqual(p, 1.0) + self.assertEqual(log10p, 0.0) + + def test_monotone_in_the_statistic(self) -> None: + K = ds.n_cells(60) + previous = 1.0 + for w in range(0, 12): + p, _ = ds.collision_pvalue(w, 400, K) + self.assertLessEqual(p, previous) + previous = p + + def test_log10_stays_finite_where_the_p_value_underflows(self) -> None: + K = ds.n_cells(1854) + p, log10p = ds.collision_pvalue(200000, 400000, K) + self.assertEqual(p, 0.0) # scipy underflows here + self.assertTrue(math.isfinite(log10p)) + self.assertLess(log10p, -300.0) + + def test_log10_agrees_with_scipy_where_both_compute(self) -> None: + K = ds.n_cells(1854) + n_rows = 1500000 + for w in [1200, 1500, 2121]: + _, log10p = ds.collision_pvalue(w, n_rows, K) + reference = math.log10(binom.sf(w - 1, math.comb(n_rows, 2), 1.0 / K)) + self.assertAlmostEqual(log10p, reference, places=2) + + def test_scipy_handles_a_realistic_number_of_pairs(self) -> None: + """Guards the pinned scipy==1.6.* against the newer one used in dev.""" + K = ds.n_cells(1854) + p, _ = ds.collision_pvalue(500, 1000000, K) + self.assertTrue(math.isfinite(p)) + self.assertGreaterEqual(p, 0.0) + self.assertLessEqual(p, 1.0) + + def test_undefined_null_is_not_a_number(self) -> None: + p, log10p = ds.collision_pvalue(5, 100, 0) + self.assertTrue(math.isnan(p)) + self.assertTrue(math.isnan(log10p)) + + +class CollisionCountsTestCase(unittest.TestCase): + def test_permuted_rows_are_the_same_trial(self) -> None: + triplets = torch.tensor([[0, 1, 2], [2, 1, 0], [1, 0, 2]]) + pairs, rows, distinct = ds.collision_counts(triplets) + self.assertEqual((pairs, rows, distinct), (3, 3, 1)) + + def test_column_order_still_distinguishes_exact_rows(self) -> None: + triplets = torch.tensor([[0, 1, 2], [2, 1, 0]]) + self.assertEqual(ds.count_duplicates(triplets), (0, 0)) + self.assertEqual(ds.collision_counts(triplets)[0], 1) + + def test_three_identical_rows_give_three_pairs(self) -> None: + triplets = torch.tensor([[0, 1, 2], [0, 1, 2], [0, 1, 2]]) + self.assertEqual(ds.collision_counts(triplets), (3, 3, 1)) + + def test_matches_brute_force(self) -> None: + torch.manual_seed(0) + triplets = torch.randint(0, 8, (50, 3)) + counts = Counter(tuple(sorted(row)) for row in triplets.tolist()) + expected_pairs = sum(math.comb(c, 2) for c in counts.values()) + expected_rows = sum(c for c in counts.values() if c > 1) + expected_distinct = sum(1 for c in counts.values() if c > 1) + self.assertEqual( + ds.collision_counts(triplets), + (expected_pairs, expected_rows, expected_distinct), + ) + + def test_degenerate_rows_are_counted(self) -> None: + triplets = torch.tensor([[5, 5, 7], [0, 1, 2], [3, 3, 3]]) + self.assertEqual(ds.count_degenerate_rows(triplets), 2) + + def test_empty_input(self) -> None: + self.assertEqual(ds.collision_counts(torch.empty(0, 3, dtype=torch.long)), (0, 0, 0)) + + +class ReportTestCase(unittest.TestCase): + @staticmethod + def clean_sample(n_rows: int, n_objects: int) -> torch.Tensor: + """Distinct triplets, so the null is satisfied exactly.""" + triplets = [] + i = j = 0 + k = 2 + while len(triplets) < n_rows: + i, j, k = (i, j, k + 1) + if k >= n_objects: + j, k = j + 1, j + 2 + if j >= n_objects - 1: + i, j, k = i + 1, i + 2, i + 3 + if k < n_objects: + triplets.append([i, j, k]) + return torch.tensor(triplets) + + def test_clean_data_is_not_significant(self) -> None: + report = ds.analyse(self.clean_sample(500, 200), n_objects=200) + self.assertEqual(report.n_colliding_pairs, 0) + self.assertEqual(report.p_value, 1.0) + self.assertFalse(report.is_significant) + + def test_self_concatenated_file_is_significant(self) -> None: + sample = self.clean_sample(500, 200) + report = ds.analyse(torch.cat((sample, sample), dim=0), n_objects=200) + self.assertEqual(report.n_colliding_pairs, 500) + self.assertTrue(report.is_significant) + self.assertEqual(report.p_value, 0.0) + self.assertTrue(math.isfinite(report.log10_p_value)) + self.assertIn("log10 p", report.summary()) + + def test_dense_regime_is_silent(self) -> None: + """Many repeats, but no more than chance predicts for M=6, N=200.""" + torch.manual_seed(0) + cells = self.clean_sample(20, 6) + sample = cells[torch.randint(0, cells.shape[0], (200,))] + report = ds.analyse(sample, n_objects=6) + self.assertGreater(report.n_duplicated_rows, 150) + self.assertFalse(report.is_significant) + + def test_undefined_null_falls_back_to_an_unconditional_warning(self) -> None: + triplets = torch.tensor([[0, 1, 2], [0, 1, 2]]) + report = ds.analyse(triplets, n_objects=2) # K = 0 + self.assertFalse(report.testable) + self.assertTrue(report.is_significant) + self.assertIn("could not be computed", report.summary()) + + def test_negative_indices_are_reported_but_not_tested(self) -> None: + triplets = torch.tensor([[-1, 1, 2], [-1, 1, 2]]) + report = ds.analyse(triplets, n_objects=10) + self.assertFalse(report.testable) + self.assertTrue(report.is_significant) + + def test_missing_n_objects_is_not_testable(self) -> None: + report = ds.analyse(torch.tensor([[0, 1, 2], [0, 1, 2]]), n_objects=None) + self.assertFalse(report.testable) + + def test_single_row_and_empty_input(self) -> None: + for triplets in [torch.tensor([[0, 1, 2]]), torch.empty(0, 3, dtype=torch.long)]: + report = ds.analyse(triplets, n_objects=10) + self.assertEqual(report.p_value, 1.0) + self.assertFalse(report.is_significant) + + def test_wrong_shape_does_not_raise(self) -> None: + report = ds.analyse(torch.tensor([0, 1, 2, 0, 1, 2]), n_objects=10) + self.assertFalse(report.testable) + self.assertFalse(report.is_significant) + + def test_a_few_extra_repeats_are_not_yet_surprising(self) -> None: + """3 collisions where chance predicts 2.4 is not evidence of anything.""" + sample = self.clean_sample(400, 60) + report = ds.analyse(torch.cat((sample, sample[:3]), dim=0), n_objects=60) + self.assertEqual(report.n_colliding_pairs, 3) + self.assertGreater(report.p_value, 0.2) + self.assertFalse(report.is_significant) + + def test_threshold_is_overridable(self) -> None: + sample = self.clean_sample(400, 60) + report = ds.analyse(torch.cat((sample, sample[:10]), dim=0), n_objects=60) + original = ds.P_VALUE_THRESHOLD + try: + ds.P_VALUE_THRESHOLD = 1e-12 + self.assertFalse(report.is_significant) + ds.P_VALUE_THRESHOLD = 0.2 + self.assertTrue(report.is_significant) + finally: + ds.P_VALUE_THRESHOLD = original diff --git a/tests/test_utils.py b/tests/test_utils.py index 5f99a74..23a2676 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -1,11 +1,13 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- +import itertools import os import shutil import torch import unittest import utils +import warnings import numpy as np import helper @@ -53,6 +55,116 @@ def test_loading(self) -> None: self.assertEqual(M, hypers['M']) +class DuplicateTripletsTestCase(unittest.TestCase): + + @staticmethod + def save_triplets(train_triplets, test_triplets) -> None: + if not os.path.exists(test_dir): + os.mkdir(test_dir) + with open(os.path.join(test_dir, train_file), 'wb') as f: + np.save(f, train_triplets) + with open(os.path.join(test_dir, test_file), 'wb') as f: + np.save(f, test_triplets) + + @staticmethod + def duplicate_rows(triplets, n_repeats: int): + """Append the first rows a second time.""" + return np.vstack((triplets, triplets[:n_repeats])) + + def test_counting(self) -> None: + triplets = torch.tensor([[0, 1, 2], [3, 4, 5], [0, 1, 2], [0, 1, 2], [6, 7, 8]]) + n_duplicated_rows, n_duplicated_triplets = utils.count_duplicates(triplets) + self.assertEqual(n_duplicated_rows, 3) + self.assertEqual(n_duplicated_triplets, 1) + + def test_counting_without_duplicates(self) -> None: + triplets = torch.tensor([[0, 1, 2], [3, 4, 5], [6, 7, 8]]) + self.assertEqual(utils.count_duplicates(triplets), (0, 0)) + + def test_permuted_rows_are_not_duplicates(self) -> None: + # the column order encodes the choice, hence a permuted row is a different trial + triplets = torch.tensor([[0, 1, 2], [2, 1, 0]]) + self.assertEqual(utils.count_duplicates(triplets), (0, 0)) + + @staticmethod + def load_and_catch_warnings(n_objects): + """Load from test_dir, returning every duplicate warning it raised.""" + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter('always') + loaded = utils.load_data( + device=device, triplets_dir=test_dir, n_objects=n_objects) + return loaded, [ + w for w in caught if issubclass(w.category, utils.DuplicateTripletsWarning) + ] + + # helper.create_triplets is unseeded, so n_objects is passed explicitly + # everywhere below: inferring it would make C(M,3), and hence every + # p-value in these tests, a random variable across runs. + + def test_warning(self) -> None: + triplets = helper.create_triplets(N=100, M=20, P=10) + train_triplets, test_triplets = helper.create_train_test_split(triplets) + n_repeats = 20 + train_triplets = self.duplicate_rows(train_triplets.numpy(), n_repeats) + self.save_triplets(train_triplets, test_triplets) + + with self.assertWarns(utils.DuplicateTripletsWarning) as cm: + loaded_train, _ = utils.load_data( + device=device, triplets_dir=test_dir, n_objects=20) + shutil.rmtree(test_dir) + + self.assertIn(f'Found {n_repeats * 2} of', str(cm.warning)) + self.assertIn(train_file.split('.')[0], str(cm.warning)) + self.assertIn('expected', str(cm.warning)) + # the warning must not drop or alter any of the loaded triplets + np.testing.assert_allclose(train_triplets, loaded_train) + + def test_no_warning(self) -> None: + triplets = helper.create_triplets(N=100, M=20, P=10) + train_triplets, test_triplets = helper.create_train_test_split(triplets) + self.save_triplets(train_triplets, test_triplets) + + _, duplicate_warnings = self.load_and_catch_warnings(n_objects=20) + shutil.rmtree(test_dir) + + self.assertEqual(len(duplicate_warnings), 0) + + def test_repeats_at_the_chance_rate_are_silent(self) -> None: + """The point of the feature: 200 draws from 20 possible triplets are + almost all repeats, and every one of them is expected.""" + rnd = np.random.RandomState(0) + cells = np.array(list(itertools.combinations(range(6), 3))) + train_triplets = cells[rnd.randint(0, len(cells), size=200)] + self.save_triplets(train_triplets, cells[:5]) + + loaded, duplicate_warnings = self.load_and_catch_warnings(n_objects=6) + shutil.rmtree(test_dir) + + # the data really is almost entirely duplicates ... + n_duplicated_rows, _ = utils.count_duplicates(loaded[0]) + self.assertGreater(n_duplicated_rows, 150) + # ... and none of it is surprising + self.assertEqual(len(duplicate_warnings), 0) + + def test_self_concatenated_file_warns_with_a_usable_p_value(self) -> None: + # 240 distinct triplets doubled puts the tail past 1e-308, which is + # exactly the case the log10 fallback exists for + triplets = helper.create_triplets(N=300, M=60, P=20) + train_triplets, test_triplets = helper.create_train_test_split(triplets) + train_triplets = train_triplets.numpy() + self.save_triplets(np.vstack((train_triplets, train_triplets)), test_triplets) + + _, duplicate_warnings = self.load_and_catch_warnings(n_objects=60) + shutil.rmtree(test_dir) + + self.assertEqual(len(duplicate_warnings), 1) + message = str(duplicate_warnings[0].message) + # p underflows float64 this far into the tail, so the message must + # carry log10 p rather than a useless 'p = 0' + self.assertIn('log10 p', message) + self.assertNotIn('p = 0.0', message) + + class CorrelationTestCase(unittest.TestCase): def test_correlation(self) -> None: diff --git a/utils.py b/utils.py index dc55490..d16ca34 100644 --- a/utils.py +++ b/utils.py @@ -5,7 +5,7 @@ import pickle from collections import defaultdict from functools import partial -from typing import Dict, Tuple +from typing import Dict, Optional, Tuple import numpy as np import pandas as pd @@ -15,9 +15,19 @@ from skimage.transform import resize from statsmodels.stats.multitest import multipletests +import duplicate_stats + Tensor = torch.Tensor Array = np.ndarray +# Re-exported so utils. keeps working. These bind the original objects +# rather than wrapping them: warn_on_duplicates reports the caller of load_data +# via stacklevel, and an extra forwarding frame would shift it by one. +DuplicateReport = duplicate_stats.DuplicateReport +DuplicateTripletsWarning = duplicate_stats.DuplicateTripletsWarning +count_duplicates = duplicate_stats.count_duplicates +warn_on_duplicates = duplicate_stats.warn_on_duplicates + def pickle_file(file: dict, out_path: str, file_name: str) -> None: with open(os.path.join(out_path, "".join((file_name, ".txt"))), "wb") as f: @@ -44,13 +54,33 @@ def load_ref_images(img_folder: str, item_names: Array) -> Array: return ref_images +def infer_nobjects(*triplets: Tensor) -> Optional[int]: + """Infer the number of objects from every triplet tensor that was loaded. + + Pooling the splits matters: deriving M from a small validation split alone + would shrink C(M, 3) and silently suppress duplicate warnings on it. + Returns None when nothing usable was loaded, since get_nobjects raises on + an empty tensor and a data check must not break the load. + """ + non_empty = [t for t in triplets if t is not None and t.shape[0] > 0] + if not non_empty: + return None + return get_nobjects(torch.cat(non_empty, dim=0)) + + def load_data( device: torch.device, triplets_dir: str, val_set: str = "test_10", inference: bool = False, + n_objects: Optional[int] = None, ) -> Tuple[Tensor]: - """Load train and test triplet datasets from disk.""" + """Load train and test triplet datasets from disk. + + Each file is checked for triplets that repeat more often than the sampling + design predicts; see duplicate_stats. Pass n_objects to pin the null's + sample space, otherwise it is inferred from the loaded data. + """ if inference: with open( os.path.join(triplets_dir, "test_triplets.npy"), "rb" @@ -60,6 +90,9 @@ def load_data( .to(device) .type(torch.LongTensor) ) + if n_objects is None: + n_objects = infer_nobjects(test_triplets) + warn_on_duplicates(test_triplets, "test_triplets", n_objects) return test_triplets try: with open(os.path.join(triplets_dir, "train_90.npy"), "rb") as train_file: @@ -84,6 +117,10 @@ def load_data( .to(device) .type(torch.LongTensor) ) + if n_objects is None: + n_objects = infer_nobjects(train_triplets, test_triplets) + warn_on_duplicates(train_triplets, "train_90", n_objects) + warn_on_duplicates(test_triplets, val_set, n_objects) return train_triplets, test_triplets