diff --git a/CHANGELOG.md b/CHANGELOG.md index dd4019d..a182b07 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -84,6 +84,35 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 though it did. Both call sites now call that function directly, which is what `loco.py` and `prepare_common.py` already do. +- **`compare_assoc_results` detects the reference test type as an enum instead of + three booleans.** `is_all_tests`, `is_score_test` and `is_lrt_test` were derived + from one sample of the reference rows, with Wald as the implicit `else` of an + if/elif chain: an enum spelled as a boolean triple. It is now `AssocTestType`, + produced by `_detect_assoc_test_type`. + + That matters because the four cases were re-derived in three separate places. + #122 merged two of them and its comment claimed they shared one dispatch, but + the SNP-count-mismatch early return still asked `is_score_test or is_all_tests` + twice to decide which optional output slots to populate, and each of the four + branches still ended in a hand-written conjunction of `.passed` attributes + repeating overlapping subsets of the same columns. + + Two tables now hold the case analysis. `_OPTIONAL_SLOTS` says which optional + result slots a test type populates, read by the early return. `_PASS_COLUMNS` + says which columns the overall verdict reads, so the four conjunctions become + one `all(...)`. LRT still omits beta and se, because it reports them as NaN by + construction and verifies both sides are all-NaN instead; that is now one named + branch rather than an inline special case inside a 240-line function. + + The per-type column selection stays a four-way branch. Those branches call + different comparison helpers with different tolerances and skip messages, so + there is nothing there to collapse. + + The file grows by 29 lines. The enum, the detector and the two tables cost more + lines than the four conjunctions they replace; what improves is that the case + analysis lives in three named places instead of being re-derived inline in + three others. + ## [7.2.0] - 2026-07-27 Minor, not major, despite the Breaking heading below. That is a deliberate diff --git a/docs/CODEMAP.md b/docs/CODEMAP.md index 4c1a62b..2db1d28 100644 --- a/docs/CODEMAP.md +++ b/docs/CODEMAP.md @@ -319,12 +319,12 @@ Tolerance-based comparison infrastructure for GEMMA parity testing. | ID | Component | Description | File:Line | |----|-----------|-------------|-----------| | 6a | `ToleranceConfig` | Per-field tolerance dataclass (strict/default/relaxed) | [tolerances.py:40](../src/jamma/validation/tolerances.py#L40) | -| 6b | `ComparisonResult` | Pass/fail with max diffs and worst location | [compare.py:21](../src/jamma/validation/compare.py#L21) | -| 6b | `AssocComparisonResult` | Per-column comparison results | [compare.py:312](../src/jamma/validation/compare.py#L312) | -| 6b | `compare_assoc_results()` | Full association comparison across test types | [compare.py:399](../src/jamma/validation/compare.py#L399) | -| 6b | `compare_kinship_matrices()` | Symmetric matrix comparison | [compare.py:144](../src/jamma/validation/compare.py#L144) | -| 6b | `load_gemma_assoc()` | Parse GEMMA `.assoc.txt` (schema-derived) | [compare.py:245](../src/jamma/validation/compare.py#L245) | -| 6b | `load_gemma_kinship()` | Parse GEMMA `.cXX.txt` | [compare.py:182](../src/jamma/validation/compare.py#L182) | +| 6b | `ComparisonResult` | Pass/fail with max diffs and worst location | [compare.py:21](../src/jamma/validation/compare.py#L22) | +| 6b | `AssocComparisonResult` | Per-column comparison results | [compare.py:312](../src/jamma/validation/compare.py#L313) | +| 6b | `compare_assoc_results()` | Full association comparison across test types | [compare.py:399](../src/jamma/validation/compare.py#L470) | +| 6b | `compare_kinship_matrices()` | Symmetric matrix comparison | [compare.py:144](../src/jamma/validation/compare.py#L145) | +| 6b | `load_gemma_assoc()` | Parse GEMMA `.assoc.txt` (schema-derived) | [compare.py:245](../src/jamma/validation/compare.py#L246) | +| 6b | `load_gemma_kinship()` | Parse GEMMA `.cXX.txt` | [compare.py:182](../src/jamma/validation/compare.py#L183) | --- @@ -582,6 +582,6 @@ Priority order: `JAMMA_BACKEND` env var -> `--backend` CLI flag -> auto (batch i | Memory estimation | [memory.py:327](../src/jamma/core/memory.py#L327) | | Threading | [threading.py:43](../src/jamma/core/threading.py#L43) | | Hardware context | [hardware.py:37](../src/jamma/core/hardware.py#L37) | -| Validation comparison | [compare.py:399](../src/jamma/validation/compare.py#L399) | +| Validation comparison | [compare.py:399](../src/jamma/validation/compare.py#L470) | | Equivalence proof | [GEMMA_EQUIVALENCE.md](GEMMA_EQUIVALENCE.md) | | Numerical equivalence bound | [GEMMA_NUMERICAL_EQUIVALENCE_BOUND.md](GEMMA_NUMERICAL_EQUIVALENCE_BOUND.md) | diff --git a/src/jamma/validation/compare.py b/src/jamma/validation/compare.py index d181c53..69d91a4 100644 --- a/src/jamma/validation/compare.py +++ b/src/jamma/validation/compare.py @@ -7,6 +7,7 @@ from __future__ import annotations from dataclasses import dataclass +from enum import StrEnum from pathlib import Path import numpy as np @@ -348,6 +349,76 @@ class AssocComparisonResult: _LAMBDA_UPPER_BOUND = 1e4 +class AssocTestType(StrEnum): + """Which GEMMA test a reference .assoc.txt came from. + + Detected from which p-value columns the reference carries, then used to + select the active comparison columns. Previously three mutually exclusive + booleans plus an implicit fourth as the ``else`` of an if/elif chain, which + meant the column selection and the overall-pass computation each had to + re-derive the same four cases and agree. + """ + + ALL = "all" + SCORE = "score" + LRT = "lrt" + WALD = "wald" + + +# Which optional output slots each test type populates. The SNP-count-mismatch +# early return reads this rather than re-deriving the test type conditions, so a +# new test type cannot populate a slot in one place and leave it None in the +# other. +_OPTIONAL_SLOTS: dict[AssocTestType, frozenset[str]] = { + AssocTestType.ALL: frozenset({"p_score", "p_lrt", "l_mle"}), + AssocTestType.SCORE: frozenset({"p_score"}), + AssocTestType.LRT: frozenset({"p_lrt", "l_mle"}), + AssocTestType.WALD: frozenset(), +} + + +# Which columns the overall pass verdict reads, per test type. LRT omits beta and +# se because it reports them as NaN by construction; compare_assoc_results checks +# both sides are all-NaN instead. +_PASS_COLUMNS: dict[AssocTestType, tuple[str, ...]] = { + AssocTestType.ALL: ( + "beta", + "se", + "af", + "p_wald", + "logl_H1", + "l_remle", + "p_score", + "p_lrt", + "l_mle", + ), + AssocTestType.SCORE: ("beta", "se", "af", "p_score"), + AssocTestType.LRT: ("af", "p_lrt", "l_mle"), + AssocTestType.WALD: ("beta", "se", "af", "p_wald", "logl_H1", "l_remle"), +} + + +def _detect_assoc_test_type(sample: list[AssocResult]) -> AssocTestType: + """Infer the test type from which p-value columns the reference carries. + + Uses a sample of the first few records rather than the first alone, so a + degenerate leading SNP with NaN columns does not decide it. + """ + if not sample: + return AssocTestType.WALD + if all( + r.p_wald is not None and r.p_lrt is not None and r.p_score is not None + for r in sample + ): + return AssocTestType.ALL + if all(r.p_wald is None for r in sample): + if all(r.p_score is not None for r in sample): + return AssocTestType.SCORE + if all(r.p_lrt is not None for r in sample): + return AssocTestType.LRT + return AssocTestType.WALD + + def _skipped_result(message: str) -> ComparisonResult: """A vacuously-passing result for a column not applicable to the test type.""" return ComparisonResult( @@ -433,27 +504,8 @@ def compare_assoc_results( if config is None: config = ToleranceConfig() - # Detect test type from reference data using majority of first N records - # (not just first record, to handle degenerate/NaN first SNP) - _SAMPLE_SIZE = min(5, len(expected)) - sample = expected[:_SAMPLE_SIZE] - - is_all_tests = len(sample) > 0 and all( - r.p_wald is not None and r.p_lrt is not None and r.p_score is not None - for r in sample - ) - is_score_test = ( - len(sample) > 0 - and not is_all_tests - and all(r.p_score is not None for r in sample) - and all(r.p_wald is None for r in sample) - ) - is_lrt_test = ( - len(sample) > 0 - and not is_all_tests - and all(r.p_lrt is not None for r in sample) - and all(r.p_wald is None for r in sample) - ) + test_type = _detect_assoc_test_type(expected[: min(5, len(expected))]) + optional_slots = _OPTIONAL_SLOTS[test_type] # Check for SNP count mismatch if len(actual) != len(expected): @@ -475,9 +527,9 @@ def compare_assoc_results( l_remle=skip_result, af=skip_result, mismatched_snps=[], - p_score=skip_result if (is_score_test or is_all_tests) else None, - p_lrt=skip_result if (is_lrt_test or is_all_tests) else None, - l_mle=skip_result if (is_lrt_test or is_all_tests) else None, + p_score=skip_result if "p_score" in optional_slots else None, + p_lrt=skip_result if "p_lrt" in optional_slots else None, + l_mle=skip_result if "l_mle" in optional_slots else None, ) # Check for mismatched SNP IDs @@ -551,89 +603,66 @@ def _lambda(field: str, *, check_upper: bool) -> ComparisonResult: no_id_mismatch = len(mismatched) == 0 - p_score_result: ComparisonResult | None = None - p_lrt_result: ComparisonResult | None = None - l_mle_result: ComparisonResult | None = None - - # Each branch computes the columns active for its test type and the overall - # pass from exactly those columns. The two must agree on which columns exist, - # so they share one dispatch rather than repeating the test-type conditions. - if is_all_tests: - pwald_result = _pvalue("p_wald", config.pvalue_rtol) - p_score_result = _pvalue("p_score", config.pvalue_rtol) - p_lrt_result = _pvalue("p_lrt", config.p_lrt_rtol) - logl_result = _logl("logl_H1 skipped (not in GEMMA -lmm 4 format)") - lambda_result = _lambda("l_remle", check_upper=False) - l_mle_result = _lambda("l_mle", check_upper=True) - all_passed = ( - beta_result.passed - and se_result.passed - and af_result.passed - and pwald_result.passed - and logl_result.passed - and lambda_result.passed - and p_score_result.passed - and p_lrt_result.passed - and l_mle_result.passed - and no_id_mismatch - ) - elif is_score_test: - p_score_result = _pvalue("p_score", config.pvalue_rtol) - pwald_result = _skipped_result("p_wald skipped (Score test)") - logl_result = _skipped_result("logl_H1 skipped (Score test)") - lambda_result = _skipped_result("l_remle skipped (Score test)") - all_passed = ( - beta_result.passed - and se_result.passed - and af_result.passed - and p_score_result.passed - and no_id_mismatch - ) - elif is_lrt_test: - p_lrt_result = _pvalue("p_lrt", config.p_lrt_rtol) - l_mle_result = _lambda("l_mle", check_upper=True) - pwald_result = _skipped_result("p_wald skipped (LRT)") - logl_result = _skipped_result("logl_H1 skipped (LRT)") - lambda_result = _skipped_result("l_remle skipped (LRT)") - # LRT has NaN beta/se by construction, so check both sides are all-NaN - # rather than the (meaningless) compare result. - beta_all_nan = bool( + # Which columns the overall pass reads, per test type. Previously each branch + # below ended in its own hand-written conjunction of `.passed` attributes, + # four overlapping copies that had to stay in step with the column selection + # immediately above them. + results: dict[str, ComparisonResult] = { + "beta": beta_result, + "se": se_result, + "af": af_result, + } + + if test_type is AssocTestType.ALL: + results["p_wald"] = _pvalue("p_wald", config.pvalue_rtol) + results["p_score"] = _pvalue("p_score", config.pvalue_rtol) + results["p_lrt"] = _pvalue("p_lrt", config.p_lrt_rtol) + results["logl_H1"] = _logl("logl_H1 skipped (not in GEMMA -lmm 4 format)") + results["l_remle"] = _lambda("l_remle", check_upper=False) + results["l_mle"] = _lambda("l_mle", check_upper=True) + elif test_type is AssocTestType.SCORE: + results["p_score"] = _pvalue("p_score", config.pvalue_rtol) + results["p_wald"] = _skipped_result("p_wald skipped (Score test)") + results["logl_H1"] = _skipped_result("logl_H1 skipped (Score test)") + results["l_remle"] = _skipped_result("l_remle skipped (Score test)") + elif test_type is AssocTestType.LRT: + results["p_lrt"] = _pvalue("p_lrt", config.p_lrt_rtol) + results["l_mle"] = _lambda("l_mle", check_upper=True) + results["p_wald"] = _skipped_result("p_wald skipped (LRT)") + results["logl_H1"] = _skipped_result("logl_H1 skipped (LRT)") + results["l_remle"] = _skipped_result("l_remle skipped (LRT)") + else: + results["p_wald"] = _pvalue("p_wald", config.pvalue_rtol) + results["logl_H1"] = _logl("logl_H1 skipped (reference missing logl_H1 column)") + results["l_remle"] = _lambda("l_remle", check_upper=False) + + # LRT reports NaN beta/se by construction, so its pass columns omit them and + # it verifies both sides are all-NaN instead of reading a meaningless compare + # result. Every other type has nothing extra to check. + if test_type is AssocTestType.LRT: + beta_se_ok = bool( np.all(np.isnan(actual_beta)) and np.all(np.isnan(expected_beta)) - ) - se_all_nan = bool(np.all(np.isnan(actual_se)) and np.all(np.isnan(expected_se))) - all_passed = ( - af_result.passed - and beta_all_nan - and se_all_nan - and p_lrt_result.passed - and l_mle_result.passed - and no_id_mismatch - ) - else: # Wald - pwald_result = _pvalue("p_wald", config.pvalue_rtol) - logl_result = _logl("logl_H1 skipped (reference missing logl_H1 column)") - lambda_result = _lambda("l_remle", check_upper=False) - all_passed = ( - beta_result.passed - and se_result.passed - and af_result.passed - and pwald_result.passed - and logl_result.passed - and lambda_result.passed - and no_id_mismatch - ) + ) and bool(np.all(np.isnan(actual_se)) and np.all(np.isnan(expected_se))) + else: + beta_se_ok = True + + all_passed = ( + all(results[column].passed for column in _PASS_COLUMNS[test_type]) + and beta_se_ok + and no_id_mismatch + ) return AssocComparisonResult( passed=all_passed, n_snps=len(actual), beta=beta_result, se=se_result, - p_wald=pwald_result, - logl_H1=logl_result, - l_remle=lambda_result, + p_wald=results["p_wald"], + logl_H1=results["logl_H1"], + l_remle=results["l_remle"], af=af_result, mismatched_snps=mismatched, - p_score=p_score_result, - p_lrt=p_lrt_result, - l_mle=l_mle_result, + p_score=results.get("p_score"), + p_lrt=results.get("p_lrt"), + l_mle=results.get("l_mle"), )