From 70a1dec249fc66c32868d4a5e12f90a0c3662b07 Mon Sep 17 00:00:00 2001 From: Nate Pope Date: Fri, 18 Sep 2026 15:44:41 -0700 Subject: [PATCH 1/9] Use pairwise-complete allele frequencies for naive r2 under missing data _pairwise_ld_core computed D = p_AB - p_A*p_B where p_AB (the joint 11-frequency) was summed over the pairwise-complete sample set (gametes valid at both sites) but p_A/p_B were each site's own marginal frequency (summed over all of that site's valid gametes). Under structured missing data those are different sample sets, so D and r2 came out biased -- in one case, a mathematically impossible r2 > 1. missing_data='exclude' never exercised this, since dropping every site with any missingness makes the marginal and joint sets trivially coincide. p_A/p_B are now conditioned on the pairwise-joint-valid set for each pair, matching the standard 2x2-table formula _tile_counts/_tile_sigma_d2 already use elsewhere in the LD module. pairwise_r2, pairwise_LD_v, and locate_unlinked (the three callers of the shared core) all pick up the fix; the tiled naive-r2 path used by zns/omega (_tile_r2_naive) is rebuilt on the same _tile_counts helper the projection estimator already uses, dropping its now-unnecessary marginal-frequency parameters. sigma_d2 is untouched: it's a different, deliberately unbiased estimator (Ragsdale & Gravel 2019), not a fix for this bug applied under another name -- it was just never affected by it, being built entirely from pairwise joint counts with no marginal-frequency vector anywhere in its own construction. --- docs/source/changelog.rst | 6 +++ pg_gpu/haplotype_matrix.py | 39 +++++++++------- pg_gpu/ld_statistics.py | 36 +++++--------- tests/test_ld_statistics_coverage.py | 70 ++++++++++++++++++++++++++++ 4 files changed, 111 insertions(+), 40 deletions(-) diff --git a/docs/source/changelog.rst b/docs/source/changelog.rst index 05f31823..144b4fbe 100644 --- a/docs/source/changelog.rst +++ b/docs/source/changelog.rst @@ -177,6 +177,12 @@ Read this section if you are comparing against older pg_gpu results. Bug fixes ~~~~~~~~~ +* ``pairwise_r2`` and the naive ``r2`` estimator behind ``zns``/``omega`` + computed the joint 11-frequency over the pairwise-complete sample but + each site's own frequency over its separate, larger marginal sample, + biasing ``D`` and ``r2`` under structured missing data with + ``missing_data='include'``. Both frequencies now come from the same + pairwise-complete sample. * Windowed Garud's H made three float64 copies of the haplotype matrix (109 GB each for 2,940 haplotypes across 4.7 million sites) and, for a Garud-only request, a transposed int8 copy on top. Each window is diff --git a/pg_gpu/haplotype_matrix.py b/pg_gpu/haplotype_matrix.py index 802cc62d..89d9a7e5 100644 --- a/pg_gpu/haplotype_matrix.py +++ b/pg_gpu/haplotype_matrix.py @@ -1763,8 +1763,8 @@ def Tajimas_D(self) -> float: def _pairwise_ld_core(self, hap_clean=None, valid_mask=None): """Shared computation for pairwise LD methods. - Computes allele frequencies, joint frequencies, and D matrix from - haplotype data, handling missing values. + Computes pairwise-complete allele frequencies, joint frequencies, and + D matrix from haplotype data, handling missing values. Parameters ---------- @@ -1778,9 +1778,11 @@ def _pairwise_ld_core(self, hap_clean=None, valid_mask=None): Returns ------- D : cupy.ndarray, shape (m, m) - Pairwise D = p_AB - p_A*p_B. - p : cupy.ndarray, shape (m,) - Per-site allele frequencies. + Pairwise D = p_AB - p_i*p_j. + p_i, p_j : cupy.ndarray, shape (m, m) + Allele frequency at site i (resp. j) restricted to the gametes + valid at both sites of the pair (m, m), not a per-site (m,) + vector. """ if self.device == 'CPU': self.transfer_to_gpu() @@ -1790,15 +1792,20 @@ def _pairwise_ld_core(self, hap_clean=None, valid_mask=None): valid_mask = (ind >= 0).astype(cp.float64) hap_clean = cp.where(ind >= 0, ind, 0).astype(cp.float64) - n_valid = cp.sum(valid_mask, axis=0).astype(cp.float64) - p = cp.where(n_valid > 0, cp.sum(hap_clean, axis=0) / n_valid, 0.0) - joint_n = valid_mask.T @ valid_mask joint_11 = hap_clean.T @ hap_clean + # p_i/p_j must share joint_n/joint_11's sample set (gametes valid at + # both sites), not each site's own larger marginal set, or D's two + # terms disagree on where their data comes from. + sum_i = hap_clean.T @ valid_mask + sum_j = sum_i.T # full index range, not a tile, so this is exact + + p_i = cp.where(joint_n > 0, sum_i / joint_n, 0.0) + p_j = cp.where(joint_n > 0, sum_j / joint_n, 0.0) p_AB = cp.where(joint_n > 0, joint_11 / joint_n, 0.0) - D = p_AB - cp.outer(p, p) - return D, p + D = p_AB - p_i * p_j + return D, p_i, p_j def pairwise_LD_v(self) -> cp.ndarray: """Pairwise linkage disequilibrium (D statistic) via matrix multiply. @@ -1809,7 +1816,7 @@ def pairwise_LD_v(self) -> cp.ndarray: from ._warnings import _warn_biallelic_only bmask = self._biallelic_mask() _warn_biallelic_only(int((~bmask).sum()), context="pairwise_LD_v") - D, _ = self._pairwise_ld_core() + D, _, _ = self._pairwise_ld_core() bad = ~bmask D[bad, :] = cp.nan D[:, bad] = cp.nan @@ -1863,9 +1870,9 @@ def pairwise_r2(self, estimator: str = 'r2') -> cp.ndarray: from ._warnings import _warn_biallelic_only bmask = self._biallelic_mask() _warn_biallelic_only(int((~bmask).sum()), context="pairwise_r2") - D, p = self._pairwise_ld_core() - denom_squared = cp.outer(p * (1 - p), p * (1 - p)) - r2 = cp.where(denom_squared > 0, (D ** 2) / denom_squared, cp.nan) + D, p_i, p_j = self._pairwise_ld_core() + denom = (p_i * (1 - p_i)) * (p_j * (1 - p_j)) + r2 = cp.where(denom > 0, (D ** 2) / denom, cp.nan) bad = ~bmask r2[bad, :] = cp.nan r2[:, bad] = cp.nan @@ -1920,11 +1927,11 @@ def locate_unlinked(self, size=100, step=20, threshold=0.1): active_idx = np.where(active)[0] + w_start active_idx_gpu = cp.asarray(active_idx) - D, p_w = self._pairwise_ld_core( + D, p_i, p_j = self._pairwise_ld_core( hap_clean[:, active_idx_gpu], valid_mask[:, active_idx_gpu], ) - denom = cp.outer(p_w * (1 - p_w), p_w * (1 - p_w)) + denom = (p_i * (1 - p_i)) * (p_j * (1 - p_j)) r2_mat = cp.where(denom > 0, (D ** 2) / denom, 0.0) cp.fill_diagonal(r2_mat, 0.0) diff --git a/pg_gpu/ld_statistics.py b/pg_gpu/ld_statistics.py index 042f1d96..68e5ccf3 100644 --- a/pg_gpu/ld_statistics.py +++ b/pg_gpu/ld_statistics.py @@ -301,13 +301,16 @@ def _tile_counts(hi, vi, hj, vj): return c1, c2, c3, c4, n -def _tile_r2_naive(hi, vi, hj, vj, pi, pqi, pj, pqj): - """Compute naive r² for a tile (frequency-based, biased).""" - joint_n = vi.T @ vj - joint_11 = hi.T @ hj - p_AB = cp.where(joint_n > 0, joint_11 / joint_n, 0.0) - D = p_AB - cp.outer(pi, pj) - denom = cp.outer(pqi, pqj) +def _tile_r2_naive(hi, vi, hj, vj): + """Compute naive r² for a tile (the classical frequency-based estimator).""" + # Built on the same pairwise-complete counts as _tile_sigma_d2, so p_i/p_j + # come from the gametes valid at both sites rather than each site's own + # (possibly larger) marginal valid set. + c1, c2, c3, c4, n = _tile_counts(hi, vi, hj, vj) + D = cp.where(n > 0, (c1 * c4 - c2 * c3) / (n * n), 0.0) + p_i = cp.where(n > 0, (c1 + c2) / n, 0.0) + p_j = cp.where(n > 0, (c1 + c3) / n, 0.0) + denom = (p_i * (1 - p_i)) * (p_j * (1 - p_j)) return cp.where(denom > 0, (D ** 2) / denom, 0.0) @@ -572,12 +575,6 @@ def _zns_tiled(mat, missing_data='include', tile_size=512, use_projection=False) total = 0.0 n_pairs = 0 - if not use_projection: - n_valid = cp.sum(valid_mask, axis=0).astype(cp.float64) - p = cp.where(n_valid > 0, - cp.sum(hap_clean, axis=0) / n_valid, 0.0) - pq = p * (1 - p) - for i0 in range(0, m, B): i1 = min(i0 + B, m) hi = hap_clean[:, i0:i1] @@ -599,9 +596,7 @@ def _zns_tiled(mat, missing_data='include', tile_size=512, use_projection=False) total += 2.0 * float(cp.sum(tile).get()) n_pairs += 2 * int(cp.sum(valid).get()) else: - r2_tile = _tile_r2_naive( - hi, vi, hj, vj, - p[i0:i1], pq[i0:i1], p[j0:j1], pq[j0:j1]) + r2_tile = _tile_r2_naive(hi, vi, hj, vj) if i0 == j0: cp.fill_diagonal(r2_tile, 0.0) total += float(cp.sum(r2_tile).get()) @@ -658,11 +653,6 @@ def _zns_from_precomputed(hap_clean, valid_mask, col_start, col_end, hc = hc[:, seg_idx] vm = vm[:, seg_idx] - if not use_projection: - n_valid = n_valid[seg_idx] - p = cp.where(n_valid > 0, cp.sum(hc, axis=0) / n_valid, 0.0) - pq = p * (1 - p) - B = tile_size total = 0.0 n_pairs = 0 @@ -688,9 +678,7 @@ def _zns_from_precomputed(hap_clean, valid_mask, col_start, col_end, total += 2.0 * float(cp.sum(tile).get()) n_pairs += 2 * int(cp.sum(valid).get()) else: - r2_tile = _tile_r2_naive( - hi, vi, hj, vj, - p[i0:i1], pq[i0:i1], p[j0:j1], pq[j0:j1]) + r2_tile = _tile_r2_naive(hi, vi, hj, vj) if i0 == j0: cp.fill_diagonal(r2_tile, 0.0) total += float(cp.sum(r2_tile).get()) diff --git a/tests/test_ld_statistics_coverage.py b/tests/test_ld_statistics_coverage.py index 7176dfee..9afc965b 100644 --- a/tests/test_ld_statistics_coverage.py +++ b/tests/test_ld_statistics_coverage.py @@ -86,6 +86,76 @@ def test_zns_naive_matches_direct_pairwise_correlation(): _agree(zns(hm, estimator="r2"), float(np.mean(r2s))) +def test_pairwise_r2_pairwise_complete_under_missing_data(): + """r2 under missing_data='include' must use one sample set -- the pair's + jointly-valid haplotypes -- for both the joint 11-frequency and each + site's own frequency, not a site's separate (larger) marginal valid set. + + Site 0 is missing on haplotypes 0-3, present (0,0,0,0,1,1,1,1,0,0,0,0) on + 4-15. Site 1 is missing on haplotypes 12-15, present + (1,1,1,1,0,0,0,0,1,1,1,1) on 0-11. Their only shared valid haplotypes are + 4-11, where the two columns are identical -- a perfect correlation, r2 + exactly 1. Each site's own marginal frequency (1/3 and 2/3, computed over + its own full valid set) differs from its frequency restricted to that + shared set (1/2 for both); using the marginal frequencies for D and the + r2 denominator gives 25/16 -- not just wrong, but impossible for a real + r2, which cannot exceed 1. + """ + hap = np.array([ + [-1, 1], [-1, 1], [-1, 1], [-1, 1], + [0, 0], [0, 0], [0, 0], [0, 0], + [1, 1], [1, 1], [1, 1], [1, 1], + [0, -1], [0, -1], [0, -1], [0, -1], + ], dtype=np.int8) + pos = np.array([100, 200], dtype=np.int64) + hm = HaplotypeMatrix(hap, pos, 0, 1000) + hm.transfer_to_gpu() + + r2 = cp.asnumpy(hm.pairwise_r2()) + assert np.isclose(r2[0, 1], 1.0, rtol=1e-9, atol=1e-12) + + # A 2-site matrix has exactly one pair, so ZnS (the tiled naive path) is + # that pair's r2 -- checks the tiled and dense paths agree. + _agree(zns(hm, estimator="r2"), 1.0) + + +def test_pairwise_r2_matches_pairwise_complete_correlation_multi_site(): + """General oracle: with several sites each missing a distinct block of + haplotypes (so every pair overlaps on a different subset), r2 under + missing_data='include' must equal the direct Pearson correlation squared + computed independently in numpy, restricted to each pair's own jointly- + valid haplotypes.""" + rng = np.random.RandomState(0) + n_hap, n_var = 24, 5 + hap = rng.randint(0, 2, size=(n_hap, n_var)).astype(np.int8) + block = n_hap // n_var + for v in range(n_var): + hap[v * block:(v + 1) * block, v] = -1 + pos = ((np.arange(n_var) + 1) * 100).astype(np.int64) + hm = HaplotypeMatrix(hap, pos, 0, 1000) + hm.transfer_to_gpu() + + r2 = cp.asnumpy(hm.pairwise_r2()) + hap_f = hap.astype(np.float64) + expected = np.zeros((n_var, n_var)) + for i in range(n_var): + for j in range(i + 1, n_var): + valid = (hap[:, i] >= 0) & (hap[:, j] >= 0) + ai, aj = hap_f[valid, i], hap_f[valid, j] + # The fix should not be exercised on a degenerate (zero-variance) + # pair -- assert the construction avoids that rather than + # silently skip it. + assert ai.std() > 0 and aj.std() > 0, (i, j) + expected[i, j] = expected[j, i] = np.corrcoef(ai, aj)[0, 1] ** 2 + + for i in range(n_var): + for j in range(i + 1, n_var): + assert np.isclose(r2[i, j], expected[i, j], rtol=1e-9, atol=1e-9), (i, j) + + iu = np.triu_indices(n_var, k=1) + _agree(zns(hm, estimator="r2"), float(expected[iu].mean())) + + @pytest.mark.parametrize("use_projection", [False, True], ids=["naive", "proj"]) def test_zns_from_precomputed_tiling_invariant(use_projection): # tile_size is an implementation detail: a small tile forces the From e132d67225027352b153916f9e4ef07941ba803b Mon Sep 17 00:00:00 2001 From: Nate Pope Date: Wed, 30 Sep 2026 09:56:31 -0700 Subject: [PATCH 2/9] Compute p_i/p_AB in place, p_j as a view, in _pairwise_ld_core Review comment on PR #312: at 8,000 sites the memory pool went from 2.66 GB to 4.2 GB from the new sum_i, p_i, p_j m x m arrays. joint_n is a self-product of the whole matrix against itself (not a tile), so p_j is exactly p_i's transpose -- a view, not a fresh computation. p_i and p_AB are now computed in place into sum_i and joint_11 via a safe (zero-substituted) denominator, since cupy's divide has no where= to mask the division directly (verified: TypeError, "Wrong arguments"). The numerator is provably 0 wherever the denominator is, so this is exact, not an approximation. Measured at 8,000 sites (40 haplotypes, 5% missing): memory pool high-water mark drops from 4.168 GB to 3.144 GB for pairwise_r2(). Values are unchanged -- existing LD test suite passes without modification. --- pg_gpu/haplotype_matrix.py | 17 +++++++++++++---- 1 file changed, 13 insertions(+), 4 deletions(-) diff --git a/pg_gpu/haplotype_matrix.py b/pg_gpu/haplotype_matrix.py index 48b0f367..18b33f0c 100644 --- a/pg_gpu/haplotype_matrix.py +++ b/pg_gpu/haplotype_matrix.py @@ -1783,11 +1783,20 @@ def _pairwise_ld_core(self, hap_clean=None, valid_mask=None): # both sites), not each site's own larger marginal set, or D's two # terms disagree on where their data comes from. sum_i = hap_clean.T @ valid_mask - sum_j = sum_i.T # full index range, not a tile, so this is exact - p_i = cp.where(joint_n > 0, sum_i / joint_n, 0.0) - p_j = cp.where(joint_n > 0, sum_j / joint_n, 0.0) - p_AB = cp.where(joint_n > 0, joint_11 / joint_n, 0.0) + # cupy's divide has no `where=`, so a safe (zero-substituted) + # denominator is the only way to divide in place. The numerator is + # provably 0 wherever joint_n is 0 (no gamete has both sites valid, + # so every gamete contributing to sum_i/joint_11 at that pair has the + # other site invalid too), so this gives the correct 0, not 0/0 -> nan. + safe_n = cp.where(joint_n > 0, joint_n, 1.0) + sum_i /= safe_n # in place: sum_i now holds p_i + joint_11 /= safe_n # in place: joint_11 now holds p_AB + p_i, p_AB = sum_i, joint_11 # aliases, not copies + # joint_n is a self-product of the whole matrix against itself (not + # a tile), so p_j is exactly p_i's transpose -- a view, not a new + # array. Must be taken after p_i is finalized above. + p_j = p_i.T D = p_AB - p_i * p_j return D, p_i, p_j From 20274ed169dba54101c1aac6cf9deb86999eaead Mon Sep 17 00:00:00 2001 From: Nate Pope Date: Wed, 30 Sep 2026 10:08:02 -0700 Subject: [PATCH 3/9] Fix the same pairwise-complete-vs-marginal bug on the genotype path Review comment on PR #312, item 1: "we need the same fix on the genotype path". _r2_matrix_diploid (the naive-r2 path for a GenotypeMatrix, reached via zns/omega(estimator='r2') and the public r2_matrix_diploid alias) computed each site's mean/variance from its own marginal valid count, while the covariance cross term was already implicitly pairwise-joint (zeroed wherever either site is invalid at a gamete). Same mismatched sample sets as the haplotype bug, same bias, same reason missing_data='exclude' never exercises it. Mean/variance are now conditioned on the pairwise-joint-valid set, using the standard pairwise-complete-observations correlation formula. A site's own dosage variance (used to NaN a monomorphic site's whole row/column, per _drop_undefined_sites's contract) stays a marginal, per-site property -- not the pairwise-conditioned one -- so a site can't be NaN for some pairs and not others. Applies the same memory discipline as the haplotype-path fix from the start: sum_j/var_j are transpose views (joint_n is a self-product of the whole matrix against itself, not a tile), and the pairwise cross/variance terms are computed in place into the base sum arrays via a safe denominator (cupy's divide has no where=). Verified against a naive pairwise-complete-obs reference implementation under structured (per-site-varying) missing data -- confirmed the old code actually disagreed with it (e.g. 0.02667 vs 0.02413 for one pair), and the fix matches exactly, including the monomorphic-site whole-row/ column NaN and no-missing-data no-op cases. Also resolves the changelog merge conflict against main (item 3 in the same review comment): both sides had added independent bullets to the Bug fixes section. --- pg_gpu/ld_statistics.py | 69 ++++++++++++++++++++-------- tests/test_ld_statistics_coverage.py | 60 +++++++++++++++++++++++- 2 files changed, 109 insertions(+), 20 deletions(-) diff --git a/pg_gpu/ld_statistics.py b/pg_gpu/ld_statistics.py index 64ec2e6f..8bb7112a 100644 --- a/pg_gpu/ld_statistics.py +++ b/pg_gpu/ld_statistics.py @@ -1039,7 +1039,7 @@ def _r2_matrix_diploid(genotype_matrix): """Compute r-squared matrix from diploid genotypes (0/1/2) on GPU. Uses genotype correlation: treats 0/1/2 as continuous dosage values, - computes Pearson correlation, then squares. + computes a pairwise-complete Pearson correlation, then squares. Parameters ---------- @@ -1063,27 +1063,58 @@ def _r2_matrix_diploid(genotype_matrix): if not isinstance(geno, cp.ndarray): geno = cp.asarray(geno) - # mask missing data: compute per-site mean from valid data only valid_mask = (geno >= 0).astype(cp.float64) geno_clean = cp.where(geno >= 0, geno, 0).astype(cp.float64) - n_valid = cp.sum(valid_mask, axis=0).astype(cp.float64) - - mean = cp.where(n_valid > 0, cp.sum(geno_clean, axis=0) / n_valid, 0.0) - - # center, zeroing out missing entries - gn = (geno_clean - mean[None, :]) * valid_mask - # variance per variant (using valid counts) - var = cp.sum(gn ** 2, axis=0) - - # r_ij = cov_ij / sqrt(var_i * var_j), applied as a rank-1 in-place - # scale so peak memory stays near the one output matrix; the NaN scale - # at a zero-variance site spreads over its whole row and column. - r2 = gn.T @ gn # (n_var, n_var) - inv = cp.where(var > 0, 1.0 / cp.sqrt(var), cp.nan) - r2 *= inv[:, None] - r2 *= inv[None, :] - r2 *= r2 + # A site's own dosage variance is a marginal (whole-sample) property, + # independent of which other site it is paired with: sites with none get + # an entirely NaN row/column, matching _drop_undefined_sites's "undefined + # entries arrive as whole rows/cols" contract. This must not use the + # pairwise-conditioned variance below, or a site could be NaN for some + # pairs and not others. + n_valid = cp.sum(valid_mask, axis=0).astype(cp.float64) + marginal_mean = cp.where(n_valid > 0, cp.sum(geno_clean, axis=0) / n_valid, 0.0) + site_defined = cp.sum(((geno_clean - marginal_mean[None, :]) * valid_mask) ** 2, + axis=0) > 0 + + # Pairwise-complete correlation: mean/variance conditioned on the gametes + # valid at both sites of a pair, not each site's own (possibly larger) + # marginal set -- same fix as HaplotypeMatrix._pairwise_ld_core, plus the + # sum-of-squares term continuous dosage needs for variance that binary + # haplotype data gets for free from the mean alone. Kept as raw + # (un-normalized) sums rather than dividing down to means/variances: the + # joint-valid count n cancels exactly in r2 = cov^2/(var_i*var_j) since + # cov/var_i/var_j are each scaled by the same n, so only one division + # per term is needed, not a second pass to re-normalize. + joint_n = valid_mask.T @ valid_mask + sum_i = geno_clean.T @ valid_mask + sum_j = sum_i.T # free transpose: self-product of the whole matrix + # against itself, not a tile -- same identity + # _pairwise_ld_core relies on + joint_11 = geno_clean.T @ geno_clean + ss_i = (geno_clean ** 2).T @ valid_mask + + # cupy's divide has no where= (verified: TypeError, "Wrong arguments"), + # so a safe (zero-substituted) denominator is the only way to divide in + # place. The numerator is provably 0 wherever joint_n is (no jointly + # valid gamete means no term in sum_i/joint_11/ss_i can be nonzero). + safe_n = cp.where(joint_n > 0, joint_n, 1.0) + joint_11 -= (sum_i * sum_j) / safe_n # in place: joint_11 now holds cov + ss_i -= (sum_i * sum_i) / safe_n # in place: ss_i now holds var_i + cov = joint_11 + var_i = ss_i + var_j = var_i.T # view, taken after var_i is finalized -- same + # symmetry argument as var_j above + + # A specific pair can still land on zero pairwise variance even when + # both sites are globally defined (the gametes jointly valid for this + # one pair happen to be constant) -- 0.0 there, matching the sentinel + # _tile_r2_naive already uses for the same situation, reserving NaN for + # the whole-row/column case above. + valid_pair = (var_i > 0) & (var_j > 0) + safe_denom = cp.where(valid_pair, var_i * var_j, 1.0) + r2 = cp.where(valid_pair, (cov * cov) / safe_denom, 0.0) + r2 = cp.where(site_defined[:, None] & site_defined[None, :], r2, cp.nan) cp.fill_diagonal(r2, 0.0) return r2 diff --git a/tests/test_ld_statistics_coverage.py b/tests/test_ld_statistics_coverage.py index 9afc965b..da365f5f 100644 --- a/tests/test_ld_statistics_coverage.py +++ b/tests/test_ld_statistics_coverage.py @@ -14,7 +14,7 @@ from pg_gpu import GenotypeMatrix, HaplotypeMatrix from pg_gpu.ld_statistics import ( - compute_ld_statistics, dd, dz, mu_ld, pi2, r, r_squared, zns, + compute_ld_statistics, dd, dz, mu_ld, omega, pi2, r, r_squared, zns, _get_pop_data, _r2_matrix_diploid, _resolve_r2_matrix, _zns_from_precomputed, ) @@ -231,6 +231,64 @@ def test_r2_matrix_diploid_zero_variance_site_is_nan(): assert np.isnan(r2[0, 1]) and np.isnan(r2[1, 0]) +def test_r2_matrix_diploid_pairwise_complete_under_missing_data(): + """Same bug as pairwise_r2's, in its continuous-dosage form: mean/variance + must come from the pair's jointly-valid individuals, not each site's own + (possibly larger) marginal valid set. + + Site 0 is missing on individuals 0-3, dosage (0,0,0,0,2,2,2,2,0,0,0,0) on + 4-15. Site 1 is missing on individuals 12-15, dosage + (2,2,2,2,0,0,0,0,2,2,2,2) on 0-11. Their only shared valid individuals are + 4-11, where the two columns are identical -- a perfect correlation, r2 + exactly 1. Each site's own marginal mean (computed over its own full + valid set) differs from its mean restricted to the shared set, so using + marginal mean/variance for centering would not give r2 == 1 here. + """ + geno = np.array([ + [-1, 2], [-1, 2], [-1, 2], [-1, 2], + [0, 0], [0, 0], [0, 0], [0, 0], + [2, 2], [2, 2], [2, 2], [2, 2], + [0, -1], [0, -1], [0, -1], [0, -1], + ], dtype=np.int8) + r2 = cp.asnumpy(_r2_matrix_diploid(geno)) + assert np.isclose(r2[0, 1], 1.0, rtol=1e-9, atol=1e-12) + + gm = GenotypeMatrix(geno, np.array([1, 2], dtype=np.int64)) + _agree(zns(gm, estimator="r2"), 1.0) + _agree(omega(gm, estimator="r2"), 0.0) # fewer than 5 sites: omega's floor + + +def test_r2_matrix_diploid_matches_pairwise_complete_correlation_multi_site(): + """General oracle: with several sites each missing a distinct block of + individuals (so every pair overlaps on a different subset), r2 must equal + the direct Pearson correlation squared computed independently in numpy, + restricted to each pair's own jointly-valid individuals.""" + rng = np.random.RandomState(1) + n_ind, n_var = 24, 5 + geno = rng.randint(0, 3, size=(n_ind, n_var)).astype(np.int8) + block = n_ind // n_var + for v in range(n_var): + geno[v * block:(v + 1) * block, v] = -1 + r2 = cp.asnumpy(_r2_matrix_diploid(geno)) + + geno_f = geno.astype(np.float64) + expected = np.zeros((n_var, n_var)) + for i in range(n_var): + for j in range(i + 1, n_var): + valid = (geno[:, i] >= 0) & (geno[:, j] >= 0) + ai, aj = geno_f[valid, i], geno_f[valid, j] + assert ai.std() > 0 and aj.std() > 0, (i, j) + expected[i, j] = expected[j, i] = np.corrcoef(ai, aj)[0, 1] ** 2 + + for i in range(n_var): + for j in range(i + 1, n_var): + assert np.isclose(r2[i, j], expected[i, j], rtol=1e-9, atol=1e-9), (i, j) + + gm = GenotypeMatrix(geno, (np.arange(n_var) + 1).astype(np.int64)) + iu = np.triu_indices(n_var, k=1) + _agree(zns(gm, estimator="r2"), float(expected[iu].mean())) + + # ── _resolve_r2_matrix (passthrough + dispatch) ──────────────────────── def test_resolve_r2_matrix_passthrough_and_dispatch(): arr = np.array([[0.0, 0.25], [0.25, 0.0]]) From c7a5c7eba1fc9295d40bb7fc407f5749540b0deb Mon Sep 17 00:00:00 2001 From: Nate Pope Date: Wed, 30 Sep 2026 10:14:15 -0700 Subject: [PATCH 4/9] Use a bool mask instead of a float64 safe denominator in _pairwise_ld_core The safe_n array (joint_n with zeros substituted, to divide in place without cupy's divide where= support) cost a full float64 (m, m) array. Dividing directly by joint_n and masking the result afterward needs only a 1-byte-per-element bool mask instead: IEEE-754 float division never traps (verified directly against numpy: cupy raises no warning for this, numpy raises two by default), it just produces 0/0 -> nan at the zero positions -- always nan here, never inf, since the numerator is provably 0 wherever joint_n is -- and the explicit zeroing overwrites every nan before anything else reads these arrays. This relies on joint_n/sum_i/joint_11 staying float64: integer divide-by-zero is undefined behavior on GPU, unlike float. Measured at 8,000 sites: 3.144 GB -> 2.888 GB (pre-PR-312 baseline is 2.632 GB). Same test results as the safe_n version. --- pg_gpu/haplotype_matrix.py | 24 ++++++++------------- pg_gpu/ld_statistics.py | 44 +++++++++++--------------------------- 2 files changed, 21 insertions(+), 47 deletions(-) diff --git a/pg_gpu/haplotype_matrix.py b/pg_gpu/haplotype_matrix.py index 18b33f0c..531e90d5 100644 --- a/pg_gpu/haplotype_matrix.py +++ b/pg_gpu/haplotype_matrix.py @@ -1779,24 +1779,18 @@ def _pairwise_ld_core(self, hap_clean=None, valid_mask=None): joint_n = valid_mask.T @ valid_mask joint_11 = hap_clean.T @ hap_clean - # p_i/p_j must share joint_n/joint_11's sample set (gametes valid at - # both sites), not each site's own larger marginal set, or D's two - # terms disagree on where their data comes from. + # p_i/p_j use the pair's joint-valid sample, not each site's own marginal one. sum_i = hap_clean.T @ valid_mask - # cupy's divide has no `where=`, so a safe (zero-substituted) - # denominator is the only way to divide in place. The numerator is - # provably 0 wherever joint_n is 0 (no gamete has both sites valid, - # so every gamete contributing to sum_i/joint_11 at that pair has the - # other site invalid too), so this gives the correct 0, not 0/0 -> nan. - safe_n = cp.where(joint_n > 0, joint_n, 1.0) - sum_i /= safe_n # in place: sum_i now holds p_i - joint_11 /= safe_n # in place: joint_11 now holds p_AB + # In-place divide by joint_n (zeros included): numerator is provably 0 + # wherever joint_n is, so this is always 0/0 -> nan, overwritten below. + zero = joint_n == 0 + sum_i /= joint_n + sum_i[zero] = 0.0 + joint_11 /= joint_n + joint_11[zero] = 0.0 p_i, p_AB = sum_i, joint_11 # aliases, not copies - # joint_n is a self-product of the whole matrix against itself (not - # a tile), so p_j is exactly p_i's transpose -- a view, not a new - # array. Must be taken after p_i is finalized above. - p_j = p_i.T + p_j = p_i.T # joint_n is symmetric, so this is exact, not a tile approximation D = p_AB - p_i * p_j return D, p_i, p_j diff --git a/pg_gpu/ld_statistics.py b/pg_gpu/ld_statistics.py index 8bb7112a..37c5a664 100644 --- a/pg_gpu/ld_statistics.py +++ b/pg_gpu/ld_statistics.py @@ -1066,51 +1066,31 @@ def _r2_matrix_diploid(genotype_matrix): valid_mask = (geno >= 0).astype(cp.float64) geno_clean = cp.where(geno >= 0, geno, 0).astype(cp.float64) - # A site's own dosage variance is a marginal (whole-sample) property, - # independent of which other site it is paired with: sites with none get - # an entirely NaN row/column, matching _drop_undefined_sites's "undefined - # entries arrive as whole rows/cols" contract. This must not use the - # pairwise-conditioned variance below, or a site could be NaN for some - # pairs and not others. + # Dosage variance is a marginal (whole-sample), not pairwise, property: + # an undefined site gets a whole NaN row/column, never a scattered one. n_valid = cp.sum(valid_mask, axis=0).astype(cp.float64) marginal_mean = cp.where(n_valid > 0, cp.sum(geno_clean, axis=0) / n_valid, 0.0) site_defined = cp.sum(((geno_clean - marginal_mean[None, :]) * valid_mask) ** 2, axis=0) > 0 - # Pairwise-complete correlation: mean/variance conditioned on the gametes - # valid at both sites of a pair, not each site's own (possibly larger) - # marginal set -- same fix as HaplotypeMatrix._pairwise_ld_core, plus the - # sum-of-squares term continuous dosage needs for variance that binary - # haplotype data gets for free from the mean alone. Kept as raw - # (un-normalized) sums rather than dividing down to means/variances: the - # joint-valid count n cancels exactly in r2 = cov^2/(var_i*var_j) since - # cov/var_i/var_j are each scaled by the same n, so only one division - # per term is needed, not a second pass to re-normalize. + # Pairwise-complete correlation, as raw sums (r2 = cov^2/(var_i*var_j) + # needs no re-normalizing, since n cancels once both are scaled by it). joint_n = valid_mask.T @ valid_mask sum_i = geno_clean.T @ valid_mask - sum_j = sum_i.T # free transpose: self-product of the whole matrix - # against itself, not a tile -- same identity - # _pairwise_ld_core relies on + sum_j = sum_i.T # free transpose, as in _pairwise_ld_core joint_11 = geno_clean.T @ geno_clean ss_i = (geno_clean ** 2).T @ valid_mask - # cupy's divide has no where= (verified: TypeError, "Wrong arguments"), - # so a safe (zero-substituted) denominator is the only way to divide in - # place. The numerator is provably 0 wherever joint_n is (no jointly - # valid gamete means no term in sum_i/joint_11/ss_i can be nonzero). + # In-place divide; numerator is provably 0 wherever joint_n is 0. safe_n = cp.where(joint_n > 0, joint_n, 1.0) - joint_11 -= (sum_i * sum_j) / safe_n # in place: joint_11 now holds cov - ss_i -= (sum_i * sum_i) / safe_n # in place: ss_i now holds var_i + joint_11 -= (sum_i * sum_j) / safe_n # now holds cov + ss_i -= (sum_i * sum_i) / safe_n # now holds var_i cov = joint_11 var_i = ss_i - var_j = var_i.T # view, taken after var_i is finalized -- same - # symmetry argument as var_j above - - # A specific pair can still land on zero pairwise variance even when - # both sites are globally defined (the gametes jointly valid for this - # one pair happen to be constant) -- 0.0 there, matching the sentinel - # _tile_r2_naive already uses for the same situation, reserving NaN for - # the whole-row/column case above. + var_j = var_i.T + + # A pair can still land on zero variance even when both sites are + # globally defined; 0.0 there, matching _tile_r2_naive's convention. valid_pair = (var_i > 0) & (var_j > 0) safe_denom = cp.where(valid_pair, var_i * var_j, 1.0) r2 = cp.where(valid_pair, (cov * cov) / safe_denom, 0.0) From 141f6d6ce1ab298360ebbdb16a8091359bc1f998 Mon Sep 17 00:00:00 2001 From: Nate Pope Date: Wed, 30 Sep 2026 12:10:04 -0700 Subject: [PATCH 5/9] Fix undefined-pair sentinel mismatch in _r2_matrix_diploid: nan, not 0.0 _r2_matrix_diploid returned 0.0 for a pair with zero jointly-valid individuals, while pairwise_r2 (the haplotype analogue) returns nan for the identical case. Root cause: the diploid path's undefined-pair handling was modeled on _tile_r2_naive's 0.0 convention, which is for a different context (zns/omega's scalar tiled aggregation) -- the right reference is pairwise_r2, since both build a full (m, m) matrix meant to be inspected directly or fed to _drop_undefined_sites, which relies on undefined entries being nan. 0.0 isn't just inconsistent, it's not mathematically derivable here: r2 = cov^2/(var_i*var_j) is 0/0 when the pairwise-restricted sample has no variance at one site, and 0/0 doesn't resolve to 0 just because the numerator is also forced to 0 by the same degeneracy -- there's no limit to take, var_i is exactly 0, not approaching it. The zero variance is a property of which individuals survived the joint-validity filter, not evidence the sites are uncorrelated. Also removes the separate marginal site_defined/n_valid/marginal_mean check: a globally invariant site has zero variance in every pairwise- restricted subsample too (a subset of constant values is still constant), so denom > 0 alone already NaNs its whole row/column -- the marginal check was redundant with it, not an independent case. Net: simpler code, three fewer (m, m) temporaries, and it now matches pairwise_r2's rule exactly. --- pg_gpu/ld_statistics.py | 21 +++++++-------------- tests/test_ld_statistics_coverage.py | 27 +++++++++++++++++++++++++++ 2 files changed, 34 insertions(+), 14 deletions(-) diff --git a/pg_gpu/ld_statistics.py b/pg_gpu/ld_statistics.py index 37c5a664..73884418 100644 --- a/pg_gpu/ld_statistics.py +++ b/pg_gpu/ld_statistics.py @@ -1049,7 +1049,9 @@ def _r2_matrix_diploid(genotype_matrix): Returns ------- r2 : cupy.ndarray, float64, shape (n_variants, n_variants) - NaN row and column at a site with no dosage variance; diagonal 0. + NaN wherever a pair's jointly-valid sample has no variance at one + site (a globally invariant site, or no individuals shared with the + other site); diagonal 0. """ from .genotype_matrix import GenotypeMatrix @@ -1066,13 +1068,6 @@ def _r2_matrix_diploid(genotype_matrix): valid_mask = (geno >= 0).astype(cp.float64) geno_clean = cp.where(geno >= 0, geno, 0).astype(cp.float64) - # Dosage variance is a marginal (whole-sample), not pairwise, property: - # an undefined site gets a whole NaN row/column, never a scattered one. - n_valid = cp.sum(valid_mask, axis=0).astype(cp.float64) - marginal_mean = cp.where(n_valid > 0, cp.sum(geno_clean, axis=0) / n_valid, 0.0) - site_defined = cp.sum(((geno_clean - marginal_mean[None, :]) * valid_mask) ** 2, - axis=0) > 0 - # Pairwise-complete correlation, as raw sums (r2 = cov^2/(var_i*var_j) # needs no re-normalizing, since n cancels once both are scaled by it). joint_n = valid_mask.T @ valid_mask @@ -1089,12 +1084,10 @@ def _r2_matrix_diploid(genotype_matrix): var_i = ss_i var_j = var_i.T - # A pair can still land on zero variance even when both sites are - # globally defined; 0.0 there, matching _tile_r2_naive's convention. - valid_pair = (var_i > 0) & (var_j > 0) - safe_denom = cp.where(valid_pair, var_i * var_j, 1.0) - r2 = cp.where(valid_pair, (cov * cov) / safe_denom, 0.0) - r2 = cp.where(site_defined[:, None] & site_defined[None, :], r2, cp.nan) + # A globally invariant site has zero variance in every pairwise-restricted + # subsample too, so this already NaNs a whole row/column, not just a pair. + denom = var_i * var_j + r2 = cp.where(denom > 0, (cov * cov) / denom, cp.nan) cp.fill_diagonal(r2, 0.0) return r2 diff --git a/tests/test_ld_statistics_coverage.py b/tests/test_ld_statistics_coverage.py index da365f5f..52e05d40 100644 --- a/tests/test_ld_statistics_coverage.py +++ b/tests/test_ld_statistics_coverage.py @@ -231,6 +231,33 @@ def test_r2_matrix_diploid_zero_variance_site_is_nan(): assert np.isnan(r2[0, 1]) and np.isnan(r2[1, 0]) +def test_r2_matrix_diploid_no_overlap_is_nan_not_zero(): + # Column 0 valid only on individuals 0-1, column 1 only on 2-3: zero + # individuals jointly valid, so this pair's correlation cannot be + # estimated at all -- nan, not a measured zero (each column is still + # individually polymorphic, so this isn't the whole-site-NaN case above). + geno = np.array([[0, -1], [1, -1], [-1, 5], [-1, 9]], dtype=np.int8) + r2 = cp.asnumpy(_r2_matrix_diploid(geno)) + assert np.isnan(r2[0, 1]) and np.isnan(r2[1, 0]) + + +def test_r2_matrix_diploid_locally_degenerate_pair_is_nan(): + # Site 0 valid on individuals 0-3 as [0, 1, 1, 1] (polymorphic overall). + # Site 1 valid on individuals 1-4 as [5, 5, 5, 9] (polymorphic overall). + # Their shared individuals are 1-3, where site 0 reads [1, 1, 1] -- + # constant in that specific jointly-valid subsample, even though neither + # site is globally monomorphic. The pair's correlation is 0/0, not 0. + geno = np.array([ + [0, -1], + [1, 5], + [1, 5], + [1, 5], + [-1, 9], + ], dtype=np.int8) + r2 = cp.asnumpy(_r2_matrix_diploid(geno)) + assert np.isnan(r2[0, 1]) and np.isnan(r2[1, 0]) + + def test_r2_matrix_diploid_pairwise_complete_under_missing_data(): """Same bug as pairwise_r2's, in its continuous-dosage form: mean/variance must come from the pair's jointly-valid individuals, not each site's own From c287f1e46b4bb8d6b11bf188ca25f5f69ee74b1d Mon Sep 17 00:00:00 2001 From: Nate Pope Date: Wed, 30 Sep 2026 20:22:06 -0700 Subject: [PATCH 6/9] Clamp joint_n in place instead of a bool mask, in _pairwise_ld_core Per PR #312 review: cp.maximum(joint_n, 1.0, out=joint_n) replaces the explicit zero-mask-then-overwrite, avoiding reliance on cupy's silent 0/0 -> nan. The numerator is provably 0 wherever joint_n was 0 (every contributing term is zeroed by invalidity at one site or the other), so dividing by the clamped value of 1 there gives the same exact 0.0 as before. Verified byte-identical to the prior bool-mask version across 20 random trials (60 haplotypes, 200 sites, 20% missing each) -- 0 mismatches in D, p_i, or p_j. Drops one (m, m) bool array. --- pg_gpu/haplotype_matrix.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/pg_gpu/haplotype_matrix.py b/pg_gpu/haplotype_matrix.py index 531e90d5..76af027f 100644 --- a/pg_gpu/haplotype_matrix.py +++ b/pg_gpu/haplotype_matrix.py @@ -1782,13 +1782,11 @@ def _pairwise_ld_core(self, hap_clean=None, valid_mask=None): # p_i/p_j use the pair's joint-valid sample, not each site's own marginal one. sum_i = hap_clean.T @ valid_mask - # In-place divide by joint_n (zeros included): numerator is provably 0 - # wherever joint_n is, so this is always 0/0 -> nan, overwritten below. - zero = joint_n == 0 + # Clamp joint_n in place before dividing: numerator is provably 0 + # wherever joint_n is 0, so the clamped divide gives exact 0 there too. + cp.maximum(joint_n, 1.0, out=joint_n) sum_i /= joint_n - sum_i[zero] = 0.0 joint_11 /= joint_n - joint_11[zero] = 0.0 p_i, p_AB = sum_i, joint_11 # aliases, not copies p_j = p_i.T # joint_n is symmetric, so this is exact, not a tile approximation From f10d3fb42fa6fd3753e5b0ba3b7d7842380d8fb3 Mon Sep 17 00:00:00 2001 From: Nate Pope Date: Wed, 30 Sep 2026 20:22:15 -0700 Subject: [PATCH 7/9] Clamp joint_n and drop safe_n in _r2_matrix_diploid; build r2 in place Per PR #312 review: mirrors the haplotype-path clamp (same provably-exact-zero argument) and drops the separate safe_n array for cp.maximum(joint_n, 1.0, out=joint_n). Builds r2 in place into cov (square, then divide by var_i and var_j in turn) instead of a separate denom array and cp.where, guarded by a boolean undefined = (var_i <= 0) | (var_j <= 0) -- variances are non-negative, so this needs only two (m, m) bool comparisons, not a float64 denom array. Measured (100 individuals, 5% missing, matching suggested benchmark): 8,000 sites peak memory 4.18 GB (pre-fix) -> 3.09 GB. Full LD test suite passes (500 passed, 10 skipped). Not bit-identical to the prior safe_n/denom version in the general (non-degenerate) case: computing cov^2 / var_i / var_j as two sequential divisions instead of cov^2 / (var_i * var_j) as one reorders the floating-point rounding (division doesn't associate). Verified empirically: ~68% of entries differ at the last bit in a random test, degenerate (undefined) pairs remain identical (both NaN). --- pg_gpu/ld_statistics.py | 24 ++++++++++++++---------- 1 file changed, 14 insertions(+), 10 deletions(-) diff --git a/pg_gpu/ld_statistics.py b/pg_gpu/ld_statistics.py index 73884418..e75da68e 100644 --- a/pg_gpu/ld_statistics.py +++ b/pg_gpu/ld_statistics.py @@ -1076,21 +1076,25 @@ def _r2_matrix_diploid(genotype_matrix): joint_11 = geno_clean.T @ geno_clean ss_i = (geno_clean ** 2).T @ valid_mask - # In-place divide; numerator is provably 0 wherever joint_n is 0. - safe_n = cp.where(joint_n > 0, joint_n, 1.0) - joint_11 -= (sum_i * sum_j) / safe_n # now holds cov - ss_i -= (sum_i * sum_i) / safe_n # now holds var_i + # Clamp joint_n in place before dividing: numerator is provably 0 + # wherever joint_n is 0, so the clamped divide gives exact 0 there too. + cp.maximum(joint_n, 1.0, out=joint_n) + joint_11 -= (sum_i * sum_j) / joint_n # now holds cov + ss_i -= (sum_i * sum_i) / joint_n # now holds var_i cov = joint_11 var_i = ss_i var_j = var_i.T - # A globally invariant site has zero variance in every pairwise-restricted - # subsample too, so this already NaNs a whole row/column, not just a pair. - denom = var_i * var_j - r2 = cp.where(denom > 0, (cov * cov) / denom, cp.nan) - cp.fill_diagonal(r2, 0.0) + # Variances are non-negative, so this is denom <= 0 without an (m, m) + # denom array; a globally invariant site NaNs a whole row/column. + undefined = (var_i <= 0) | (var_j <= 0) + cov *= cov # now holds cov^2 + cov /= var_i + cov /= var_j # now holds r2 = cov^2 / (var_i * var_j) + cov[undefined] = cp.nan + cp.fill_diagonal(cov, 0.0) - return r2 + return cov # Keep old names as aliases for backward compat From b2e54413355e15f3bc7f286a647e80661037d758 Mon Sep 17 00:00:00 2001 From: Nate Pope Date: Wed, 30 Sep 2026 22:03:07 -0700 Subject: [PATCH 8/9] Exclude scattered undefined pairs in omega/zns, not just dropped sites Per PR #312 review: pairwise-complete frequencies can leave a single pair undefined (zero jointly-valid sample, or a jointly-valid subsample that happens to be locally constant) while both its sites are otherwise fine. _drop_undefined_sites only drops a site when its whole row is NaN, so that scattered NaN reached omega's cp.cumsum, poisoning the prefix sum and collapsing results to a spurious 0.0 (demonstrated: haplotype omega 1.786 -> 0.0, genotype omega 1.495 -> 0.0). zns's generic path was only mildly biased (its nansum already dropped the pair from the numerator, but m * (m - 1) in the denominator didn't shrink to match: 0.0358 -> 0.0304). Fix, in omega(): fill NaN to 0 before cumsum, and replace the closed-form pair-count formulas with a real prefix-summed "defined-pair" count built the same way as the r^2 prefix sum -- the cross_sum > 0 divide-by-zero guard stays, now doing only its original job since cross_sum is never NaN anymore. Fix, in zns()'s generic tail: pair count from the finite mask itself, not m * (m - 1). Also fixes the same bug in zns's separate tiled naive-r2 path (_zns_tiled / _zns_from_precomputed / _tile_r2_naive), which is what zns() on a HaplotypeMatrix with estimator='r2' actually dispatches to and so is exact demonstrated regression -- the generic-path fix above doesn't reach it. _tile_r2_naive now returns a validity mask alongside r2 (mirroring _tile_sigma_d2's existing pattern), and both tiled functions accumulate n_pairs from it instead of assuming m * (m - 1); the existing parity test for this path (test_pairwise_r2_matches_pairwise_complete_correlation_multi_site) explicitly excluded the degenerate case it needed to catch this, so a new test closes that gap. Verified against brute-force references (not the prior buggy code): hand-built 5x5 matrix for omega/zns's matrix path, hand-built HaplotypeMatrix for the tiled path. Full LD test suite passes (519 passed, 10 skipped). --- docs/source/changelog.rst | 8 ++++ pg_gpu/ld_statistics.py | 70 +++++++++++++++++++--------- tests/test_ld_statistics_coverage.py | 57 ++++++++++++++++++++++ 3 files changed, 112 insertions(+), 23 deletions(-) diff --git a/docs/source/changelog.rst b/docs/source/changelog.rst index 4400fa6b..aecbaa8e 100644 --- a/docs/source/changelog.rst +++ b/docs/source/changelog.rst @@ -210,6 +210,14 @@ Bug fixes continuous form -- mean/variance from each site's own marginal valid set instead of the pair's jointly-valid individuals -- and is fixed the same way. +* ``omega`` and ``zns`` mishandled a single pair left undefined by the + fix above (both sites otherwise fine, just no -- or locally + degenerate -- jointly-valid sample): the NaN propagated through + ``omega``'s prefix sum and collapsed results to ``0.0``, and both + statistics divided by a closed-form pair count that assumed every + pair was defined. Both now count only their actually-defined pairs, + in the full-matrix path and in ``zns``'s tiled naive-``r2`` path + (``_zns_tiled``/``_zns_from_precomputed``). * ``pbs`` counted a population's own diversity at a site where another population in the pair had no data, which pulled that pair's FST down; a single such site could change the sign of PBS. PBS now uses, for all diff --git a/pg_gpu/ld_statistics.py b/pg_gpu/ld_statistics.py index e75da68e..600b0e6c 100644 --- a/pg_gpu/ld_statistics.py +++ b/pg_gpu/ld_statistics.py @@ -302,7 +302,12 @@ def _tile_counts(hi, vi, hj, vj): def _tile_r2_naive(hi, vi, hj, vj): - """Compute naive r² for a tile (the classical frequency-based estimator).""" + """Compute naive r² for a tile (the classical frequency-based estimator). + + Returns r2 and a parallel validity mask: an undefined pair (denom <= 0) + gets 0.0 in r2, not NaN, so a caller that must exclude it from a mean + needs the mask, not isnan. + """ # Built on the same pairwise-complete counts as _tile_sigma_d2, so p_i/p_j # come from the gametes valid at both sites rather than each site's own # (possibly larger) marginal valid set. @@ -311,7 +316,9 @@ def _tile_r2_naive(hi, vi, hj, vj): p_i = cp.where(n > 0, (c1 + c2) / n, 0.0) p_j = cp.where(n > 0, (c1 + c3) / n, 0.0) denom = (p_i * (1 - p_i)) * (p_j * (1 - p_j)) - return cp.where(denom > 0, (D ** 2) / denom, 0.0) + valid = denom > 0 + r2 = cp.where(valid, (D ** 2) / denom, 0.0) + return r2, valid def _tile_sigma_d2(hi, vi, hj, vj): @@ -596,16 +603,17 @@ def _zns_tiled(mat, missing_data='include', tile_size=512, use_projection=False) total += 2.0 * float(cp.sum(tile).get()) n_pairs += 2 * int(cp.sum(valid).get()) else: - r2_tile = _tile_r2_naive(hi, vi, hj, vj) + r2_tile, valid_tile = _tile_r2_naive(hi, vi, hj, vj) if i0 == j0: cp.fill_diagonal(r2_tile, 0.0) + cp.fill_diagonal(valid_tile, False) total += float(cp.sum(r2_tile).get()) + n_pairs += int(cp.sum(valid_tile).get()) else: total += 2.0 * float(cp.sum(r2_tile).get()) + n_pairs += 2 * int(cp.sum(valid_tile).get()) - if use_projection: - return total / n_pairs if n_pairs > 0 else 0.0 - return total / (m * (m - 1)) + return total / n_pairs if n_pairs > 0 else 0.0 def _zns_from_precomputed(hap_clean, valid_mask, col_start, col_end, @@ -678,16 +686,17 @@ def _zns_from_precomputed(hap_clean, valid_mask, col_start, col_end, total += 2.0 * float(cp.sum(tile).get()) n_pairs += 2 * int(cp.sum(valid).get()) else: - r2_tile = _tile_r2_naive(hi, vi, hj, vj) + r2_tile, valid_tile = _tile_r2_naive(hi, vi, hj, vj) if i0 == j0: cp.fill_diagonal(r2_tile, 0.0) + cp.fill_diagonal(valid_tile, False) total += float(cp.sum(r2_tile).get()) + n_pairs += int(cp.sum(valid_tile).get()) else: total += 2.0 * float(cp.sum(r2_tile).get()) + n_pairs += 2 * int(cp.sum(valid_tile).get()) - if use_projection: - return total / n_pairs if n_pairs > 0 else 0.0 - return total / (m * (m - 1)) + return total / n_pairs if n_pairs > 0 else 0.0 def _drop_undefined_sites(r2_matrix): @@ -697,8 +706,11 @@ def _drop_undefined_sites(r2_matrix): r^2 is undefined: ``pairwise_r2`` for monomorphic and multiallelic sites, ``_r2_matrix_diploid`` for sites with no dosage variance. So excluding undefined pairs is the same as dropping those sites. - No-op on a finite matrix. Assumes undefined entries arrive as whole - rows/cols (all this package produces); a scattered NaN would propagate. + No-op on a finite matrix. Only drops whole rows/cols; a single + undefined pair within an otherwise-defined site is left in place. + ``zns`` and ``omega`` separately count only their defined pairs, so + a scattered NaN doesn't bias them; other callers of this function + would need the same care. """ finite = ~cp.isnan(r2_matrix) cp.fill_diagonal(finite, False) @@ -778,8 +790,12 @@ def zns(r2_matrix_or_matrix, missing_data='include', estimator='auto'): m = int(cp.any(finite, axis=1).sum()) if m < 2: return 0.0 + # Pair count from the finite mask itself, not m * (m - 1): a scattered + # undefined pair (both sites otherwise fine) doesn't drop a whole site + # from m, so the mean must exclude it from the denominator too. + n_pairs = int(finite.sum()) total = cp.nansum(r2_matrix) - cp.nansum(cp.diag(r2_matrix)) - return float((total / (m * (m - 1))).get()) + return float((total / n_pairs).get()) def _zns_biallelic(hm, missing_data='include', estimator='auto'): @@ -899,31 +915,40 @@ def omega(r2_matrix_or_matrix, missing_data='include', estimator='auto'): # work with upper triangle only (i < j), matching diploSHIC r2 = cp.triu(r2_matrix, k=1) - - # 2D prefix sums on upper triangle + undefined = cp.isnan(r2) + # A scattered undefined pair (both its sites otherwise fine) must not + # propagate through cumsum, and the pair counts below can't assume + # full density the way a closed-form formula would. defined is built + # from r2_matrix directly (not r2) so triu's own zeroed-out lower + # triangle/diagonal isn't miscounted as defined pairs. + r2 = cp.where(undefined, 0.0, r2) + defined = cp.triu((~cp.isnan(r2_matrix)).astype(cp.int64), k=1) + + # 2D prefix sums on upper triangle, of both the r^2 values and which + # pairs are defined S = cp.cumsum(cp.cumsum(r2, axis=0), axis=1) + C = cp.cumsum(cp.cumsum(defined, axis=0), axis=1) # partition points l = 3..m-2 (matching diploSHIC) l_vals = cp.arange(3, m - 1) # left block: upper triangle pairs (i,j) with i < j < l left_sum = S[l_vals - 1, l_vals - 1] + left_count = C[l_vals - 1, l_vals - 1] # total upper triangle sum total_upper = S[m - 1, m - 1] + total_count = C[m - 1, m - 1] # cross block: pairs (i,j) with i < l and j >= l cross_sum = S[l_vals - 1, m - 1] - left_sum + cross_count = C[l_vals - 1, m - 1] - left_count # right block: pairs (i,j) with i >= l and j > i (upper triangle of right block) right_sum = total_upper - left_sum - cross_sum - # pair counts (upper triangle only) - n_left = l_vals * (l_vals - 1) // 2 - n_right = (m - l_vals) * (m - l_vals - 1) // 2 - n_cross = l_vals * (m - l_vals) - - n_within = n_left + n_right + n_within = total_count - cross_count + n_cross = cross_count within_sum = left_sum + right_sum valid = (n_within > 0) & (n_cross > 0) & (cross_sum > 0) @@ -1085,8 +1110,7 @@ def _r2_matrix_diploid(genotype_matrix): var_i = ss_i var_j = var_i.T - # Variances are non-negative, so this is denom <= 0 without an (m, m) - # denom array; a globally invariant site NaNs a whole row/column. + # Variances are non-negative, so this is denom <= 0 without an (m, m) denom array. undefined = (var_i <= 0) | (var_j <= 0) cov *= cov # now holds cov^2 cov /= var_i diff --git a/tests/test_ld_statistics_coverage.py b/tests/test_ld_statistics_coverage.py index 52e05d40..20bd397c 100644 --- a/tests/test_ld_statistics_coverage.py +++ b/tests/test_ld_statistics_coverage.py @@ -156,6 +156,36 @@ def test_pairwise_r2_matches_pairwise_complete_correlation_multi_site(): _agree(zns(hm, estimator="r2"), float(expected[iu].mean())) +def test_zns_tiled_excludes_undefined_pair_unlike_closed_form_count(): + """Counterpart to the multi-site test above, which deliberately avoids a + degenerate pair: here sites 0 and 1 share no valid haplotype at all (an + undefined pair, not a measured zero), while every other pair is defined. + _zns_tiled must exclude it from both the sum and the pair count, not + divide by the closed-form m*(m-1) that assumes every pair is defined.""" + hap = np.array([ + [0, -1, 0, 1], + [1, -1, 1, 1], + [0, -1, 0, 0], + [1, -1, 1, 0], + [-1, 1, 0, 1], + [-1, 0, 1, 1], + [-1, 1, 0, 0], + [-1, 0, 1, 0], + ], dtype=np.int8) + pos = np.array([100, 200, 300, 400], dtype=np.int64) + hm = HaplotypeMatrix(hap, pos, 0, 1000) + hm.transfer_to_gpu() + + r2 = cp.asnumpy(hm.pairwise_r2()) + assert np.isnan(r2[0, 1]) and np.isnan(r2[1, 0]) + + finite = ~np.isnan(r2) + np.fill_diagonal(finite, False) + expected = (np.nansum(r2) - np.nansum(np.diag(r2))) / finite.sum() + + _agree(zns(hm, estimator="r2"), float(expected)) + + @pytest.mark.parametrize("use_projection", [False, True], ids=["naive", "proj"]) def test_zns_from_precomputed_tiling_invariant(use_projection): # tile_size is an implementation detail: a small tile forces the @@ -258,6 +288,33 @@ def test_r2_matrix_diploid_locally_degenerate_pair_is_nan(): assert np.isnan(r2[0, 1]) and np.isnan(r2[1, 0]) +def _hand_r2_matrix_one_scattered_undefined_pair(): + """5x5 r2 matrix, every pair defined except (1, 3), to check that omega + and zns exclude a single undefined pair from their own sums and pair + counts rather than treating it as a dropped site (_drop_undefined_sites + only drops whole rows) or letting it propagate through a cumulative sum.""" + upper = { + (0, 1): 0.1, (0, 2): 0.2, (0, 3): 0.3, (0, 4): 0.4, + (1, 2): 0.5, (1, 3): None, (1, 4): 0.6, + (2, 3): 0.7, (2, 4): 0.8, + (3, 4): 0.9, + } + r2 = np.zeros((5, 5)) + for (i, j), v in upper.items(): + r2[i, j] = r2[j, i] = np.nan if v is None else v + return r2 + + +def test_omega_excludes_single_undefined_pair_not_whole_sites(): + r2 = _hand_r2_matrix_one_scattered_undefined_pair() + assert omega(cp.asarray(r2)) == pytest.approx(0.7589285714285713) + + +def test_zns_excludes_single_undefined_pair_from_pair_count(): + r2 = _hand_r2_matrix_one_scattered_undefined_pair() + assert zns(cp.asarray(r2)) == pytest.approx(0.5) + + def test_r2_matrix_diploid_pairwise_complete_under_missing_data(): """Same bug as pairwise_r2's, in its continuous-dosage form: mean/variance must come from the pair's jointly-valid individuals, not each site's own From 9159d5fc8dde7ac7ead3d9763d582be9f2e7b263 Mon Sep 17 00:00:00 2001 From: Nate Pope Date: Thu, 1 Oct 2026 08:52:11 -0700 Subject: [PATCH 9/9] Build pairwise_r2's tail in place, dropping the separate denom array HaplotypeMatrix.pairwise_r2's tail built denom = (p_i*(1-p_i))*(p_j*(1-p_j)) and r2 = cp.where(denom > 0, (D**2)/denom, nan) as fresh (m, m) arrays on top of the 3 persistent arrays _pairwise_ld_core already returns. Mirrors the clamp/guard pattern used elsewhere in this PR: p_i *= (1 - p_i) turns p_i into var_i (p_j, a transpose view of the same buffer, becomes var_j for free), a bool (var_i <= 0) | (var_j <= 0) guard replaces the float64 denom array, and D is squared and divided in place to become r2 with no separate output array. Measured peak GPU memory (100 haplotypes, 4,000/8,000 sites, 5% missing), before vs. after on this branch, as a multiple of one (m, m) float64 array: 6.08x/6.04x -> 5.08x/5.04x. A full single-array's worth saved, confirmed empirically rather than assumed from the temporary count. There is no equivalent issue for GenotypeMatrix. --- pg_gpu/haplotype_matrix.py | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/pg_gpu/haplotype_matrix.py b/pg_gpu/haplotype_matrix.py index 76af027f..0b22d318 100644 --- a/pg_gpu/haplotype_matrix.py +++ b/pg_gpu/haplotype_matrix.py @@ -1857,8 +1857,14 @@ def pairwise_r2(self, estimator: str = 'r2') -> cp.ndarray: bmask = self._biallelic_mask() _warn_biallelic_only(int((~bmask).sum()), context="pairwise_r2") D, p_i, p_j = self._pairwise_ld_core() - denom = (p_i * (1 - p_i)) * (p_j * (1 - p_j)) - r2 = cp.where(denom > 0, (D ** 2) / denom, cp.nan) + p_i *= (1 - p_i) # now var_i; p_j (a transpose view of p_i) becomes var_j for free + var_i, var_j = p_i, p_j + undefined = (var_i <= 0) | (var_j <= 0) + D *= D + D /= var_i + D /= var_j # D now holds r2, in place + D[undefined] = cp.nan + r2 = D bad = ~bmask r2[bad, :] = cp.nan r2[:, bad] = cp.nan