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
319 changes: 319 additions & 0 deletions src/vsparse/_norm_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -386,6 +386,205 @@ def _fill_normalized_major_is_row(
out[i, c] = (_g(scaled, g_code) - col_mean[c]) * col_post_scale[c]


# -- squared-norm statistics (whole array + per-condition slices) -----------
#
# Both reductions use the same ``||A_norm||_F^2 = sum(Delta^2) - 2 sum(Delta *
# offset[col]) + n_rows * sum(offset^2)`` expansion the *sparse-plus-external-
# means* code path in ``parafac2.utils.calc_norm_sq``/``calc_slice_norms``
# uses, just carried out against ``Delta`` (this view's own uncentered,
# scaled sparse term -- see :mod:`vsparse._vcs_matmul`) instead of a plain
# scipy ``data``/``means`` pair, since ``offset = col_post_scale * col_mean``
# is exactly the external ``means`` that convention expects.
#
# For VCSR (major=rows), a whole major slice belongs to exactly one row, so
# ``sum(Delta^2)``/``sum(Delta * offset)`` are pure scalar (norm_sq) or
# per-condition (slice_norms) reductions with no cross-thread write conflict:
# numba recognizes plain ``+=`` accumulation in a ``prange`` loop as a
# reduction. For VCSC (major=cols), a column's nonzeros span many rows/
# conditions, so ``slice_norms`` needs the same thread-chunked scatter as
# :func:`_gstats_col_sums_vcs`; ``norm_sq`` only ever needs a scalar, so it
# stays reduction-only even there.


@numba.njit(cache=True, parallel=True)
def _norm_sq_terms_major_is_row(
major_ptr, values, value_ptr, indices, row_scale, gene_scale, col_mean, col_post_scale, g_code
):
n_major = major_ptr.shape[0] - 1
total_sq = 0.0
total_cross = 0.0
for i in numba.prange(n_major): # ty: ignore[not-iterable]
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]):
col = indices[k]
gs = gene_scale[col]
if gs <= 0.0:
continue
s = col_post_scale[col]
delta = s * _g(v / rs / gs, g_code)
total_sq += delta * delta
total_cross += delta * col_mean[col] * s
return total_sq, total_cross


@numba.njit(cache=True, parallel=True)
def _norm_sq_terms_major_is_col(
major_ptr, values, value_ptr, indices, row_scale, gene_scale, col_mean, col_post_scale, g_code
):
n_major = major_ptr.shape[0] - 1
total_sq = 0.0
total_cross = 0.0
for j in numba.prange(n_major): # ty: ignore[not-iterable]
gs = gene_scale[j]
if gs <= 0.0:
continue
s = col_post_scale[j]
offset = col_mean[j] * s
col_sq = 0.0
col_sum = 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]):
row = indices[k]
delta = s * _g(v / row_scale[row] / gs, g_code)
col_sq += delta * delta
col_sum += delta
total_sq += col_sq
total_cross += col_sum * offset
return total_sq, total_cross


@numba.njit(cache=True, parallel=True)
def _slice_norm_terms_major_is_row(
major_ptr,
values,
value_ptr,
indices,
row_scale,
gene_scale,
col_mean,
col_post_scale,
g_code,
condition_idxs,
n_cond,
nthreads,
):
n_major = major_ptr.shape[0] - 1
chunk = (n_major + nthreads - 1) // nthreads
partial_sq = np.zeros((nthreads, n_cond), dtype=np.float64)
partial_cross = np.zeros((nthreads, n_cond), dtype=np.float64)
for t in numba.prange(nthreads): # ty: ignore[not-iterable]
start = t * chunk
end = min(n_major, start + chunk)
loc_sq = partial_sq[t]
loc_cross = partial_cross[t]
for i in range(start, end):
rs = row_scale[i]
cond = condition_idxs[i]
row_sq = 0.0
row_cross = 0.0
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]):
col = indices[k]
gs = gene_scale[col]
if gs <= 0.0:
continue
s = col_post_scale[col]
delta = s * _g(v / rs / gs, g_code)
row_sq += delta * delta
row_cross += delta * col_mean[col] * s
loc_sq[cond] += row_sq
loc_cross[cond] += row_cross
return partial_sq.sum(axis=0), partial_cross.sum(axis=0)


@numba.njit(cache=True, parallel=True)
def _slice_norm_terms_major_is_col(
major_ptr,
values,
value_ptr,
indices,
row_scale,
gene_scale,
col_mean,
col_post_scale,
g_code,
condition_idxs,
n_cond,
nthreads,
):
n_major = major_ptr.shape[0] - 1
chunk = (n_major + nthreads - 1) // nthreads
partial_sq = np.zeros((nthreads, n_cond), dtype=np.float64)
partial_cross = np.zeros((nthreads, n_cond), dtype=np.float64)
for t in numba.prange(nthreads): # ty: ignore[not-iterable]
start = t * chunk
end = min(n_major, start + chunk)
loc_sq = partial_sq[t]
loc_cross = partial_cross[t]
for j in range(start, end):
gs = gene_scale[j]
if gs <= 0.0:
continue
s = col_post_scale[j]
cm = col_mean[j]
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]):
row = indices[k]
cond = condition_idxs[row]
delta = s * _g(v / row_scale[row] / gs, g_code)
loc_sq[cond] += delta * delta
loc_cross[cond] += delta * cm * s
return partial_sq.sum(axis=0), partial_cross.sum(axis=0)


# -- uncentered sparse materialization ---------------------------------------
#
# Unlike ``toarray``'s dense fill, these write only the structural nonzeros
# (``Delta``, see :mod:`vsparse._vcs_matmul`) into a flat ``nnz``-length
# buffer, in exactly the position ``indices`` already gives each one -- so
# ``(data, indices, value_ptr[major_ptr])`` is directly a valid scipy CSR/CSC
# triple with the same sparsity pattern as the underlying raw array. The
# per-gene mean correction (:attr:`NormalizedViewBase.means`) is left for the
# caller to subtract externally, matching the convention
# ``parafac2.utils.calc_norm_sq``/``calc_W`` already use for a sparse ``X``
# plus a separate ``means`` vector.


@numba.njit(cache=True, parallel=True)
def _materialize_delta_major_is_row(
major_ptr, values, value_ptr, indices, row_scale, gene_scale, col_post_scale, g_code, data
):
n_major = major_ptr.shape[0] - 1
for i in numba.prange(n_major): # ty: ignore[not-iterable]
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]):
col = indices[k]
gs = gene_scale[col]
data[k] = col_post_scale[col] * _g(v / rs / gs, g_code) if gs > 0.0 else 0.0


@numba.njit(cache=True, parallel=True)
def _materialize_delta_major_is_col(
major_ptr, values, value_ptr, indices, row_scale, gene_scale, col_post_scale, g_code, data
):
n_major = major_ptr.shape[0] - 1
for j in numba.prange(n_major): # ty: ignore[not-iterable]
gs = gene_scale[j]
s = col_post_scale[j]
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]):
row = indices[k]
data[k] = s * _g(v / row_scale[row] / gs, g_code) if gs > 0.0 else 0.0


def _prep_key(key: Any) -> Any:
"""Turn a bare int into a length-1 list, so fancy indexing never drops that axis."""
if isinstance(key, int | np.integer):
Expand Down Expand Up @@ -616,6 +815,17 @@ def s(self) -> np.ndarray:
"""Per-gene post-scale."""
return self.col_post_scale

@property
def means(self) -> np.ndarray:
"""Per-gene mean-correction vector, ``col_post_scale * col_mean``.

This is the ``means`` a caller following the ``parafac2``-style
convention (a sparse/uncentered matrix plus a separate per-column
``means`` vector) should subtract externally: ``self.toarray() ==
self.to_scipy_sparse().toarray() - self.means``.
"""
return self.col_mean * self.col_post_scale

@property
def shape(self) -> tuple[int, int]:
return self._arr.shape
Expand Down Expand Up @@ -664,6 +874,115 @@ def toarray(self) -> np.ndarray:
)
return out

def to_scipy_sparse(self) -> Any:
"""The uncentered, scaled sparse ``Delta`` term, as a real scipy sparse array.

Same sparsity pattern as the underlying raw array (a ``csr_array``
for a VCSR-backed view, ``csc_array`` for VCSC), with :attr:`means`
left to subtract externally -- see :attr:`means`. Useful for handing
this view to code (such as ``parafac2``'s CuPy/MLX GPU backends)
that only knows how to move a plain NumPy/SciPy array onto a device,
rather than this view's own ``__matmul__``/``__rmatmul__``.
"""
import scipy.sparse as sp

arr = self._arr
data = np.empty(arr.nnz, dtype=np.float64)
if self._format == "csc":
_materialize_delta_major_is_col(
arr.major_ptr,
arr.values,
arr.value_ptr,
arr.indices,
self.row_scale,
self.gene_scale,
self.col_post_scale,
self.recipe.g_code,
data,
)
ctor = sp.csc_array
else:
_materialize_delta_major_is_row(
arr.major_ptr,
arr.values,
arr.value_ptr,
arr.indices,
self.row_scale,
self.gene_scale,
self.col_post_scale,
self.recipe.g_code,
data,
)
ctor = sp.csr_array

indptr = arr.value_ptr[arr.major_ptr]
return ctor((data, arr.indices, indptr), shape=self.shape)

def norm_sq(self) -> float:
"""Squared Frobenius norm of the full normalized matrix, in ``O(nnz + n_cols)``."""
arr = self._arr
args = (
arr.major_ptr,
arr.values,
arr.value_ptr,
arr.indices,
self.row_scale,
self.gene_scale,
self.col_mean,
self.col_post_scale,
self.recipe.g_code,
)
if self._format == "csc":
total_sq, total_cross = _norm_sq_terms_major_is_col(*args)
else:
total_sq, total_cross = _norm_sq_terms_major_is_row(*args)

n_rows = self.shape[0]
offset_sq_sum = float(np.sum(self.means**2))
return float(total_sq - 2.0 * total_cross + n_rows * offset_sq_sum)

def slice_norms(self, condition_idxs: Any, n_cond: int) -> np.ndarray:
"""Per-condition Frobenius norm of the normalized matrix's rows.

Parameters
----------
condition_idxs : array-like of int
Condition index (in ``[0, n_cond)``) for each row.
n_cond : int
The total number of conditions.

Returns
-------
np.ndarray
Length-``n_cond`` array of each condition's rows' Frobenius norm.
"""
idxs = np.asarray(condition_idxs, dtype=np.int64)
arr = self._arr
nthreads = numba.get_num_threads()
args = (
arr.major_ptr,
arr.values,
arr.value_ptr,
arr.indices,
self.row_scale,
self.gene_scale,
self.col_mean,
self.col_post_scale,
self.recipe.g_code,
idxs,
n_cond,
nthreads,
)
if self._format == "csc":
total_sq, total_cross = _slice_norm_terms_major_is_col(*args)
else:
total_sq, total_cross = _slice_norm_terms_major_is_row(*args)

counts = np.bincount(idxs, minlength=n_cond).astype(np.float64)
offset_sq_sum = float(np.sum(self.means**2))
sq = total_sq - 2.0 * total_cross + counts * offset_sq_sum
return np.sqrt(np.clip(sq, 0.0, None))

# -- selection ---------------------------------------------------------------

def select(self, rows: Any = slice(None), cols: Any = slice(None)) -> Any:
Expand Down
Loading
Loading