Skip to content
Merged
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
138 changes: 107 additions & 31 deletions src/vsparse/_norm_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -234,20 +234,20 @@ def _g_np(x: np.ndarray, g_code: int) -> np.ndarray:

@numba.njit(cache=True, parallel=True)
def _column_stats_major_is_col(
major_ptr, values, value_ptr, indices, row_scale, need_b, need_gstats, g_code
major_ptr, values, value_ptr, indices, row_scale, need_b, need_gstats, g_code, n_rows
):
"""Per-column ``gsum`` (raw material for ``b``) and sum/sum-of-squares of ``g(scaled)``.
"""Per-column ``gsum`` (raw material for ``b``) and the ``g(scaled)`` variance inputs.

Fused into one pass per column (rather than two separate dispatches):
unlike the VCSR scatter passes below, a VCSC column's ``gsum`` depends
only on that column's own nonzeros, so it's already final by the time
the second (``g``-transform) loop over the same nonzeros needs it --
no need to wait for every other column to finish first.
Sums cover the stored entries only; :func:`_finish_variance` folds in the
implicit zeros. A VCSC column's statistics depend only on its own
nonzeros, so all three loops fuse into one pass per column.
"""
n_major = major_ptr.shape[0] - 1
gsum = np.ones(n_major, dtype=np.float64)
col_sum = np.zeros(n_major, dtype=np.float64)
col_sumsq = np.zeros(n_major, dtype=np.float64)
col_m2 = np.zeros(n_major, dtype=np.float64)
col_corr = np.zeros(n_major, dtype=np.float64)
col_nnz = np.zeros(n_major, dtype=np.float64)
for j in numba.prange(n_major): # ty: ignore[not-iterable]
gs = 1.0
if need_b:
Expand All @@ -259,17 +259,30 @@ def _column_stats_major_is_col(
gsum[j] = gs
if need_gstats and gs > 0.0:
s0 = 0.0
s1 = 0.0
count = 0.0
for u in range(major_ptr[j], major_ptr[j + 1]):
v = values[u]
for k in range(value_ptr[u], value_ptr[u + 1]):
scaled = v / row_scale[indices[k]] / gs
gy = _g(scaled, g_code)
s0 += gy
s1 += gy * gy
s0 += _g(v / row_scale[indices[k]] / gs, g_code)
count += 1.0
col_sum[j] = s0
col_sumsq[j] = s1
return gsum, col_sum, col_sumsq
col_nnz[j] = count
# Deviations about the mean, not `sumsq - mean ** 2`: for a
# column with no spread every deviation is 0 exactly, where the
# latter subtracts two numbers of size `mean ** 2` and keeps only
# their rounding error.
mean = s0 / n_rows if n_rows > 0 else 0.0
m2 = 0.0
corr = 0.0
for u in range(major_ptr[j], major_ptr[j + 1]):
v = values[u]
for k in range(value_ptr[u], value_ptr[u + 1]):
d = _g(v / row_scale[indices[k]] / gs, g_code) - mean
m2 += d * d
corr += d
col_m2[j] = m2
col_corr[j] = corr
return gsum, col_sum, col_m2, col_corr, col_nnz


# -- statistics: major=rows -- scatter-add passes ----------------------------
Expand Down Expand Up @@ -297,15 +310,47 @@ def _scaled_col_sums_vcs(major_ptr, values, value_ptr, indices, row_scale, n_col
def _gstats_col_sums_vcs(
major_ptr, values, value_ptr, indices, row_scale, gene_scale, g_code, n_cols, nthreads
):
"""Per-column sum of ``g(scaled)`` over stored entries, and how many there were."""
n_major = major_ptr.shape[0] - 1
chunk = (n_major + nthreads - 1) // nthreads
partial_sum = np.zeros((nthreads, n_cols), dtype=np.float64)
partial_sumsq = np.zeros((nthreads, n_cols), dtype=np.float64)
partial_nnz = np.zeros((nthreads, n_cols), dtype=np.float64)
for t in numba.prange(nthreads): # ty: ignore[not-iterable]
start = t * chunk
end = min(n_major, start + chunk)
local_sum = partial_sum[t]
local_sumsq = partial_sumsq[t]
local_nnz = partial_nnz[t]
for i in range(start, end):
rs = row_scale[i]
for u in range(major_ptr[i], major_ptr[i + 1]):
v = values[u]
for k in range(value_ptr[u], value_ptr[u + 1]):
c = indices[k]
gs = gene_scale[c]
if gs > 0.0:
local_sum[c] += _g(v / rs / gs, g_code)
local_nnz[c] += 1.0
return partial_sum.sum(axis=0), partial_nnz.sum(axis=0)


@numba.njit(cache=True, parallel=True)
def _gstats_col_deviations_vcs(
major_ptr, values, value_ptr, indices, row_scale, gene_scale, g_code, col_mean, n_cols, nthreads
):
"""Per-column ``sum(y - mean)`` and ``sum((y - mean) ** 2)`` over stored entries.

A second walk of the values: on a row-major layout a column's mean is not
final until every row has been scattered into it.
"""
n_major = major_ptr.shape[0] - 1
chunk = (n_major + nthreads - 1) // nthreads
partial_m2 = np.zeros((nthreads, n_cols), dtype=np.float64)
partial_corr = np.zeros((nthreads, n_cols), dtype=np.float64)
for t in numba.prange(nthreads): # ty: ignore[not-iterable]
start = t * chunk
end = min(n_major, start + chunk)
local_m2 = partial_m2[t]
local_corr = partial_corr[t]
for i in range(start, end):
rs = row_scale[i]
for u in range(major_ptr[i], major_ptr[i + 1]):
Expand All @@ -314,11 +359,26 @@ def _gstats_col_sums_vcs(
c = indices[k]
gs = gene_scale[c]
if gs > 0.0:
scaled = v / rs / gs
gy = _g(scaled, g_code)
local_sum[c] += gy
local_sumsq[c] += gy * gy
return partial_sum.sum(axis=0), partial_sumsq.sum(axis=0)
d = _g(v / rs / gs, g_code) - col_mean[c]
local_m2[c] += d * d
local_corr[c] += d
return partial_m2.sum(axis=0), partial_corr.sum(axis=0)


def _finish_variance(col_sum, m2_stored, corr_stored, col_nnz, n_rows):
"""Per-column mean and variance, given either layout's stored-only sums."""
if n_rows <= 0:
zeros = np.zeros_like(col_sum)
return zeros, zeros
mean = col_sum / n_rows
# Implicit zeros sit at `g(0) == 0`, so each deviates by exactly `-mean`
# and their contribution is closed-form rather than iterated.
n_zero = n_rows - col_nnz
m2 = m2_stored + n_zero * mean**2
corr = corr_stored - n_zero * mean
# `corr ** 2 / n` is the corrected two-pass term (Chan, Golub & LeVeque):
# the deviations are about the computed mean, not the exact one.
return mean, np.clip((m2 - corr**2 / n_rows) / n_rows, 0.0, None)


# -- full materialization -----------------------------------------------------
Expand Down Expand Up @@ -608,6 +668,17 @@ def _compute_row_scale(arr: Any, recipe: Recipe) -> np.ndarray:
return row_scale


def _is_constant_column(variance: np.ndarray, mean: np.ndarray, n_rows: int) -> np.ndarray:
"""Whether each column's variance is indistinguishable from zero."""
# A constant column has no unit-variance scaling to be put on, and
# dividing by whatever the arithmetic left behind amplifies rounding noise
# into an arbitrary O(1) value. Bound from scikit-learn's
# `_is_constant_feature` (Chan, Golub & LeVeque), sized to the error
# `_finish_variance` can leave.
eps = np.finfo(np.float64).eps
return variance <= n_rows * eps * variance + (n_rows * mean * eps) ** 2


class NormalizedViewBase:
"""Shared implementation for the normalized VCSC/VCSR views.

Expand Down Expand Up @@ -658,7 +729,7 @@ def __init__(
if self._format == "csc":
# One fused pass per column for both -- see _column_stats_major_is_col.
if need_b or need_gstats:
gene_scale, col_sum, col_sumsq = _column_stats_major_is_col(
gene_scale, col_sum, col_m2, col_corr, col_nnz = _column_stats_major_is_col(
arr.major_ptr,
arr.values,
arr.value_ptr,
Expand All @@ -667,10 +738,11 @@ def __init__(
need_b,
need_gstats,
self.recipe.g_code,
n_rows,
)
else:
gene_scale = np.ones(n_cols, dtype=np.float64)
col_sum = col_sumsq = np.zeros(n_cols, dtype=np.float64)
col_sum = col_m2 = col_corr = col_nnz = np.zeros(n_cols, dtype=np.float64)
else:
# VCSR can't fuse these: gene_scale[c] isn't final until every row
# has been scattered into it, so the g-transform pass has to wait
Expand All @@ -684,31 +756,35 @@ def __init__(
gene_scale = np.ones(n_cols, dtype=np.float64)
if need_gstats:
nthreads = numba.get_num_threads()
col_sum, col_sumsq = _gstats_col_sums_vcs(
kernel_args = (
arr.major_ptr,
arr.values,
arr.value_ptr,
indices,
row_scale,
gene_scale,
self.recipe.g_code,
)
col_sum, col_nnz = _gstats_col_sums_vcs(*kernel_args, n_cols, nthreads)
col_m2, col_corr = _gstats_col_deviations_vcs(
*kernel_args,
col_sum / n_rows if n_rows > 0 else np.zeros(n_cols, dtype=np.float64),
n_cols,
nthreads,
)
else:
col_sum = col_sumsq = np.zeros(n_cols, dtype=np.float64)
col_sum = col_m2 = col_corr = col_nnz = np.zeros(n_cols, dtype=np.float64)
self.gene_scale = gene_scale

if need_gstats:
mean = col_sum / n_rows if n_rows > 0 else np.zeros(n_cols, dtype=np.float64)
variance = np.clip(
col_sumsq / n_rows - mean**2 if n_rows > 0 else np.zeros(n_cols), 0.0, None
)
mean, variance = _finish_variance(col_sum, col_m2, col_corr, col_nnz, n_rows)
std = np.sqrt(variance)
col_mean = mean if self.recipe.center else np.zeros(n_cols, dtype=np.float64)
if self.recipe.post_scale:
with np.errstate(divide="ignore", invalid="ignore"):
col_post_scale = np.where(std > 0.0, 1.0 / std, 1.0)
col_post_scale = np.where(
_is_constant_column(variance, mean, n_rows), 1.0, 1.0 / std
)
else:
col_post_scale = np.ones(n_cols, dtype=np.float64)
else:
Expand Down
8 changes: 7 additions & 1 deletion tests/test_property_normalization.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,7 +83,13 @@ def _recipe_reference(dense: np.ndarray, recipe: str) -> np.ndarray:
if recipe in ("scanpy", "pearson"):
std = g.std(axis=0)
with np.errstate(divide="ignore", invalid="ignore"):
s = np.where(std > 0, 1.0 / std, 1.0)
# Constant columns have no unit-variance scaling; same bound as
# the library's `_is_constant_column`, over numpy's own std.
eps = np.finfo(np.float64).eps
n = dense.shape[0]
var = std**2
constant = var <= n * eps * var + (n * g.mean(axis=0) * eps) ** 2
s = np.where(constant, 1.0, 1.0 / std)
else:
s = np.ones(dense.shape[1])

Expand Down
59 changes: 57 additions & 2 deletions tests/test_vcs_norm_recipes.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
import scipy.sparse as sp

from vsparse import RECIPES, Recipe, VCSCAnnData, VCSCArray, VCSRArray
from vsparse._norm_common import NORM_CACHE_MAXSIZE
from vsparse._norm_common import NORM_CACHE_MAXSIZE, _is_constant_column


@pytest.fixture(params=[VCSCArray, VCSRArray])
Expand Down Expand Up @@ -66,7 +66,13 @@ def _reference(dense: np.ndarray, recipe: str) -> np.ndarray:
if recipe in ("scanpy", "pearson"):
std = g.std(axis=0)
with np.errstate(divide="ignore", invalid="ignore"):
s = np.where(std > 0, 1.0 / std, 1.0)
# Constant columns have no unit-variance scaling; same bound as
# the library's `_is_constant_column`, over numpy's own std.
eps = np.finfo(np.float64).eps
n = dense.shape[0]
var = std**2
constant = var <= n * eps * var + (n * g.mean(axis=0) * eps) ** 2
s = np.where(constant, 1.0, 1.0 / std)
else:
s = np.ones(dense.shape[1])

Expand Down Expand Up @@ -356,3 +362,52 @@ def test_anndata_cache_does_not_pin_a_dropped_view():
assert ref() is None
# Still reusable -- from obs/varm/uns if not from the retained statistics.
assert adata.normalized("scanpy", recalculate=False) is not None


@pytest.mark.parametrize("n_rows", [2, 5, 6, 17, 64, 501])
@pytest.mark.parametrize("value", [1.0, 7.0, 9999.0])
@pytest.mark.parametrize("recipe", ["scanpy", "pearson"])
def test_a_column_with_no_variance_centers_to_zero(vcls, n_rows, value, recipe):
"""Depth normalization flattens a one-column matrix, so centering leaves 0."""
dense = np.full((n_rows, 1), value)
v = vcls.from_scipy(_scipy_for(vcls, dense))
out = v.normalized(recipe).toarray()
np.testing.assert_allclose(out, np.zeros_like(dense), atol=1e-12)


def test_a_constant_column_does_not_suppress_its_neighbours(vcls):
"""Zeroing a no-variance column must not touch the columns beside it."""
rng = np.random.default_rng(0)
dense = rng.integers(1, 50, size=(40, 5)).astype(float)
dense[:, 2] = 4.0
v = vcls.from_scipy(_scipy_for(vcls, dense))
out = v.normalized("pearson").toarray()
varying = np.delete(out, 2, axis=1)
assert np.abs(varying).max() > 0.5


@pytest.mark.parametrize("n_rows", [6, 10_000, 1_300_000])
def test_small_but_real_variance_is_not_called_constant(n_rows):
"""A coefficient of variation of 1e-06 is signal, not noise, at any scale here."""
mean = np.array([9.21])
std = 1e-6 * mean[0]
assert not _is_constant_column(np.array([std**2]), mean, n_rows)[0]


@pytest.mark.parametrize("n_rows", [6, 10_000, 1_300_000])
def test_constant_detection_still_catches_a_flat_column(n_rows):
"""Variance down at the arithmetic's own noise floor is treated as zero."""
mean = np.array([9.21])
noise = n_rows * np.finfo(np.float64).eps * mean[0]
assert _is_constant_column(np.array([(noise * 0.1) ** 2]), mean, n_rows)[0]


def test_both_layouts_compute_the_same_variance(dense):
"""Both layouts must produce the same statistics."""
if dense.sum() == 0:
pytest.skip("all-zero matrix: median row total is 0")
for recipe in ("scanpy", "pearson", "parafac2"):
r = VCSRArray.from_scipy(sp.csr_array(dense)).normalized(recipe)
c = VCSCArray.from_scipy(sp.csc_array(dense)).normalized(recipe)
np.testing.assert_allclose(r.col_post_scale, c.col_post_scale, rtol=1e-12)
np.testing.assert_allclose(r.col_mean, c.col_mean, rtol=1e-12, atol=1e-15)
Loading