mirror of
https://github.com/chanzuckerberg/cellxgene.git
synced 2026-09-30 08:08:12 +08:00
fix for incorrect stats computation in diff exp t-test (#2318)
* 2211 fixes * lint * lint * add missing test and bug found by test * change terminology for count distribution * update scanpy requirement * update scanpy requirement
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
import numpy as np
|
||||
from scipy import sparse, stats
|
||||
from backend.common.constants import XApproxDistribution
|
||||
|
||||
|
||||
def diffexp_ttest(adaptor, maskA, maskB, top_n=8, diffexp_lfc_cutoff=0.01):
|
||||
@@ -7,7 +8,7 @@ def diffexp_ttest(adaptor, maskA, maskB, top_n=8, diffexp_lfc_cutoff=0.01):
|
||||
Return differential expression statistics for top N variables.
|
||||
|
||||
Algorithm:
|
||||
- compute log fold change (log2(meanA/meanB))
|
||||
- compute fold change
|
||||
- compute Welch's t-test statistic and pvalue (w/ Bonferroni correction)
|
||||
- return top N abs(logfoldchange) where lfc > diffexp_lfc_cutoff
|
||||
|
||||
@@ -26,21 +27,24 @@ def diffexp_ttest(adaptor, maskA, maskB, top_n=8, diffexp_lfc_cutoff=0.01):
|
||||
:param top_n: number of variables to return stats for
|
||||
:param diffexp_lfc_cutoff: minimum
|
||||
absolute value returning [ varindex, logfoldchange, pval, pval_adj ] for top N genes
|
||||
:return: for top N genes, {"positive": for top N genes, [ varindex, logfoldchange, pval, pval_adj ], "negative": for top N genes, [ varindex, logfoldchange, pval, pval_adj ]}
|
||||
:return: for top N genes, {"positive": for top N genes, [ varindex, foldchange, pval, pval_adj ], "negative": for top N genes, [ varindex, foldchange, pval, pval_adj ]}
|
||||
"""
|
||||
|
||||
X_approx_distribution = adaptor.get_X_approx_distribution()
|
||||
dataA = adaptor.get_X_array(maskA, None)
|
||||
dataB = adaptor.get_X_array(maskB, None)
|
||||
|
||||
# mean, variance, N - calculate for both selections
|
||||
meanA, vA, nA = mean_var_n(dataA)
|
||||
meanB, vB, nB = mean_var_n(dataB)
|
||||
meanA, vA, nA = mean_var_n(dataA, X_approx_distribution)
|
||||
meanB, vB, nB = mean_var_n(dataB, X_approx_distribution)
|
||||
res = diffexp_ttest_from_mean_var(meanA, vA, nA, meanB, vB, nB, top_n, diffexp_lfc_cutoff)
|
||||
|
||||
return res
|
||||
|
||||
|
||||
def diffexp_ttest_from_mean_var(meanA, varA, nA, meanB, varB, nB, top_n, diffexp_lfc_cutoff):
|
||||
# IMPORTANT NOTE: this code assumes the data is normally distributed and/or already logged.
|
||||
|
||||
n_var = meanA.shape[0]
|
||||
top_n = min(top_n, n_var)
|
||||
|
||||
@@ -64,15 +68,15 @@ def diffexp_ttest_from_mean_var(meanA, varA, nA, meanB, varB, nB, top_n, diffexp
|
||||
pvals_adj = pvals * n_var
|
||||
pvals_adj[pvals_adj > 1] = 1 # cap adjusted p-value at 1
|
||||
|
||||
# logfoldchanges: log2(meanA / meanB)
|
||||
logfoldchanges = np.log2(np.abs((meanA + 1e-9) / (meanB + 1e-9)))
|
||||
# log fold change. The data is normally distributed/logged, so just subtract the means.
|
||||
logfoldchanges = meanA - meanB
|
||||
|
||||
stats_to_sort = tscores
|
||||
# find all with lfc > cutoff
|
||||
lfc_above_cutoff_idx = np.nonzero(np.abs(logfoldchanges) > diffexp_lfc_cutoff)[0]
|
||||
|
||||
# derive sort order
|
||||
if lfc_above_cutoff_idx.shape[0] > top_n*2:
|
||||
if lfc_above_cutoff_idx.shape[0] > top_n * 2:
|
||||
# partition top N
|
||||
rel_t_partition = np.argpartition(stats_to_sort[lfc_above_cutoff_idx], (top_n, -top_n))
|
||||
rel_t_partition_top_n = np.concatenate((rel_t_partition[-top_n:], rel_t_partition[:top_n]))
|
||||
@@ -95,16 +99,21 @@ def diffexp_ttest_from_mean_var(meanA, varA, nA, meanB, varB, nB, top_n, diffexp
|
||||
pvals_adj_top_n = pvals_adj[sort_order]
|
||||
|
||||
# varIndex, logfoldchange, pval, pval_adj
|
||||
result = {"positive": [[sort_order[i], logfoldchanges_top_n[i], pvals_top_n[i], pvals_adj_top_n[i]] for i in
|
||||
range(top_n)],
|
||||
"negative": [[sort_order[i], logfoldchanges_top_n[i], pvals_top_n[i], pvals_adj_top_n[i]] for i in
|
||||
range(-1, -1 - top_n, -1)], }
|
||||
result = {
|
||||
"positive": [
|
||||
[sort_order[i], logfoldchanges_top_n[i], pvals_top_n[i], pvals_adj_top_n[i]] for i in range(top_n)
|
||||
],
|
||||
"negative": [
|
||||
[sort_order[i], logfoldchanges_top_n[i], pvals_top_n[i], pvals_adj_top_n[i]]
|
||||
for i in range(-1, -1 - top_n, -1)
|
||||
],
|
||||
}
|
||||
|
||||
return result
|
||||
|
||||
|
||||
# Convenience function which handles sparse data
|
||||
def mean_var_n(X):
|
||||
def mean_var_n(X, X_approx_distribution=XApproxDistribution.NORMAL):
|
||||
"""
|
||||
Two-pass variance calculation. Numerically (more) stable
|
||||
than naive methods (and same method used by numpy.var())
|
||||
@@ -122,16 +131,27 @@ def mean_var_n(X):
|
||||
with np.errstate(divide="call", invalid="call", call=fp_err_set):
|
||||
n = X.shape[0]
|
||||
if sparse.issparse(X):
|
||||
if X_approx_distribution == XApproxDistribution.COUNT:
|
||||
X = X.log1p()
|
||||
mean = X.mean(axis=0).A1
|
||||
dfm = X - mean
|
||||
sumsq = np.sum(np.multiply(dfm, dfm), axis=0).A1
|
||||
v = sumsq / (n - 1)
|
||||
else:
|
||||
if X_approx_distribution == XApproxDistribution.COUNT:
|
||||
X = np.log1p(X)
|
||||
mean = X.mean(axis=0)
|
||||
dfm = X - mean
|
||||
sumsq = np.sum(np.multiply(dfm, dfm), axis=0)
|
||||
v = sumsq / (n - 1)
|
||||
|
||||
# AnnData does not guarantee that operations on a view of X will
|
||||
# return an ndarray, so force the cast if it wasn't done for us.
|
||||
if type(mean) is not np.ndarray:
|
||||
mean = mean.toarray()
|
||||
if type(v) is not np.ndarray:
|
||||
v = v.toarray()
|
||||
|
||||
if fp_err_occurred:
|
||||
mean[np.isfinite(mean) == False] = 0 # noqa: E712
|
||||
v[np.isfinite(v) == False] = 0 # noqa: E712
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
import numba
|
||||
import concurrent.futures
|
||||
import numpy as np
|
||||
from scipy import sparse
|
||||
from backend.common.constants import XApproxDistribution
|
||||
|
||||
|
||||
@numba.njit(fastmath=True, error_model="numpy", nogil=True)
|
||||
def min_max(arr):
|
||||
"""Return (min, max) values for the ndarray."""
|
||||
n = arr.size
|
||||
odd = n % 2
|
||||
if not odd:
|
||||
n -= 1
|
||||
max_val = min_val = arr[0]
|
||||
i = 1
|
||||
while i < n:
|
||||
x = arr[i]
|
||||
y = arr[i + 1]
|
||||
if x > y:
|
||||
x, y = y, x
|
||||
min_val = min(x, min_val)
|
||||
max_val = max(y, max_val)
|
||||
i += 2
|
||||
if not odd:
|
||||
x = arr[n]
|
||||
min_val = min(x, min_val)
|
||||
max_val = max(x, max_val)
|
||||
return min_val, max_val
|
||||
|
||||
|
||||
def estimate_approximate_distribution(X) -> XApproxDistribution:
|
||||
"""
|
||||
Estimate the distribution (normal, count) of the X matrix.
|
||||
|
||||
Currently this is based upon the assumption that scRNA-seq data is
|
||||
exponentially distributed in its raw (count) form, and when logged,
|
||||
any (max-min) range in excess of 24 is implies tens of millions of
|
||||
observations of a single feature and so is extremely unlikely.
|
||||
"""
|
||||
if sparse.isspmatrix_csc(X) or sparse.isspmatrix_csr(X):
|
||||
Xdata = X.data
|
||||
elif type(X) is np.ndarray:
|
||||
Xdata = X.reshape(
|
||||
X.size,
|
||||
)
|
||||
else:
|
||||
raise TypeError(f"Unsupported matrix type: {str(type(X))}")
|
||||
|
||||
CHUNKSIZE = 1 << 24
|
||||
if Xdata.size > CHUNKSIZE:
|
||||
min_val = max_val = Xdata[0]
|
||||
with concurrent.futures.ThreadPoolExecutor() as tp:
|
||||
for (_min, _max) in tp.map(min_max, [Xdata[i : i + CHUNKSIZE] for i in range(0, Xdata.size, CHUNKSIZE)]):
|
||||
min_val = min(_min, min_val)
|
||||
max_val = max(_max, max_val)
|
||||
|
||||
else:
|
||||
min_val, max_val = min_max(Xdata)
|
||||
|
||||
excess_range = (max_val - min_val) > 24
|
||||
return XApproxDistribution.COUNT if excess_range else XApproxDistribution.NORMAL
|
||||
@@ -24,6 +24,11 @@ class DiffExpMode(AugmentedEnum):
|
||||
VAR_FILTER = "varFilter"
|
||||
|
||||
|
||||
class XApproxDistribution(AugmentedEnum):
|
||||
NORMAL = "normal"
|
||||
COUNT = "count"
|
||||
|
||||
|
||||
JSON_NaN_to_num_warning_msg = "JSON encoding failure - please verify all data are finite values (no NaN or Infinities)"
|
||||
REACTIVE_LIMIT = 1_000_000
|
||||
|
||||
|
||||
Reference in New Issue
Block a user