Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions docs/source/changelog.rst
Original file line number Diff line number Diff line change
Expand Up @@ -199,6 +199,25 @@ 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. The diploid dosage-correlation path
(``_r2_matrix_diploid``, reached by ``zns``/``omega`` on a
``GenotypeMatrix`` with ``estimator='r2'``) had the same bug in its
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
Expand Down
48 changes: 31 additions & 17 deletions pg_gpu/haplotype_matrix.py
Original file line number Diff line number Diff line change
Expand Up @@ -1748,8 +1748,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
----------
Expand All @@ -1763,9 +1763,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()
Expand All @@ -1775,15 +1777,21 @@ 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_AB = cp.where(joint_n > 0, joint_11 / joint_n, 0.0)
# 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

# 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
joint_11 /= joint_n
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

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.
Expand All @@ -1794,7 +1802,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
Expand Down Expand Up @@ -1848,9 +1856,15 @@ 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()
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
Expand Down Expand Up @@ -1905,11 +1919,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)

Expand Down
150 changes: 85 additions & 65 deletions pg_gpu/ld_statistics.py
Original file line number Diff line number Diff line change
Expand Up @@ -301,14 +301,24 @@ 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)
return cp.where(denom > 0, (D ** 2) / denom, 0.0)
def _tile_r2_naive(hi, vi, hj, vj):
"""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.
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))
valid = denom > 0
r2 = cp.where(valid, (D ** 2) / denom, 0.0)
return r2, valid


def _tile_sigma_d2(hi, vi, hj, vj):
Expand Down Expand Up @@ -572,12 +582,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]
Expand All @@ -599,18 +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,
p[i0:i1], pq[i0:i1], p[j0:j1], pq[j0:j1])
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,
Expand Down Expand Up @@ -658,11 +661,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
Expand All @@ -688,18 +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,
p[i0:i1], pq[i0:i1], p[j0:j1], pq[j0:j1])
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):
Expand All @@ -709,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)
Expand Down Expand Up @@ -790,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'):
Expand Down Expand Up @@ -911,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)
Expand Down Expand Up @@ -1051,7 +1064,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
----------
Expand All @@ -1061,7 +1074,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

Expand All @@ -1075,30 +1090,35 @@ 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
cp.fill_diagonal(r2, 0.0)

return r2
# 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, as in _pairwise_ld_core
joint_11 = geno_clean.T @ geno_clean
ss_i = (geno_clean ** 2).T @ valid_mask

# 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

# 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
cov /= var_j # now holds r2 = cov^2 / (var_i * var_j)
cov[undefined] = cp.nan
cp.fill_diagonal(cov, 0.0)

return cov


# Keep old names as aliases for backward compat
Expand Down
Loading
Loading