Skip to content
Draft
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
41 changes: 33 additions & 8 deletions src/vsparse/_anndata_class.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,13 @@

from vsparse import _compression, _io
from vsparse._base import VCSCArray, VCSRArray, _VCSBase
from vsparse._norm_common import DEFAULT_RECIPE, Recipe, _NormCache, resolve_recipe
from vsparse._norm_common import (
DEFAULT_RECIPE,
NormalizedViewBase,
Recipe,
_NormCache,
resolve_recipe,
)
from vsparse._vcs_norm import VCSCArrayNormalized, VCSRArrayNormalized

if TYPE_CHECKING:
Expand All @@ -23,6 +29,12 @@
__all__ = ["VCSCAnnData"]

_VCS_TYPES = (VCSCArray, VCSRArray)
#: Types `X` may hold. A normalized view is included because it is the whole
#: point of `normalized()`: a lazy, `.select`-able object that streaming
#: consumers can restrict without materializing. Excluding it forced every
#: caller through `to_scipy_sparse()`, which materializes the entire matrix --
#: 42 GB on a cohort-scale dataset -- purely to have somewhere to put it.
_X_TYPES = (VCSCArray, VCSRArray, NormalizedViewBase)
_AnyVCS = _VCSBase
_DF_KEYS = ("obs", "var")
_MAPPING_KEYS = ("obsm", "varm", "obsp", "varp", "layers", "uns")
Expand Down Expand Up @@ -56,6 +68,12 @@ def _subset_1d(v: Any, idx: Any) -> Any:
def _subset_2d(v: Any, oidx: Any, vidx: Any) -> Any:
if v is None:
return None
if isinstance(v, NormalizedViewBase):
# `select(..., recalculate=False)` keeps this view's statistics instead
# of renormalizing the subset on its own -- the same semantics as
# slicing an already-normalized dense/sparse matrix, which is what a
# caller subsetting an AnnData expects.
return v.select(oidx, vidx, recalculate=False)
if isinstance(v, _VCS_TYPES):
result = v[oidx, vidx]
# VCSCArray/VCSRArray.__getitem__ only converts to a plain scipy
Expand All @@ -80,11 +98,18 @@ def _copy_value(v: Any) -> Any:
return None if v is None else v.copy()


def _check_vcs_type(value: Any, name: str) -> None:
if value is not None and not isinstance(value, _VCS_TYPES):
def _check_vcs_type(value: Any, name: str, allowed: tuple[type, ...] = _VCS_TYPES) -> None:
"""Reject anything that is not a stored array, or (for ``X``) a view of one.

``raw_X`` keeps the narrower set: it holds the raw counts by definition, so
a normalized view is not a thing it can meaningfully be.
"""
if value is not None and not isinstance(value, allowed):
extra = ", or a normalized view of one" if NormalizedViewBase in allowed else ""
build = ", or .normalized(...)" if NormalizedViewBase in allowed else ""
raise TypeError(
f"{name} must be a VCSCArray or VCSRArray, got {type(value).__name__}. "
f"Build one with VCSCArray.from_scipy(...) or vsparse.from_anndata(...)."
f"{name} must be a VCSCArray or VCSRArray{extra}, got {type(value).__name__}. "
f"Build one with VCSCArray.from_scipy(...) or vsparse.from_anndata(...){build}."
)


Expand Down Expand Up @@ -120,7 +145,7 @@ def __init__(
raw_X: _AnyVCS | None = None,
**kwargs: Any,
) -> None:
_check_vcs_type(X, "X")
_check_vcs_type(X, "X", _X_TYPES)
_check_vcs_type(raw_X, "raw_X")
if "raw" in kwargs:
raise TypeError(
Expand All @@ -146,12 +171,12 @@ def X(self) -> _AnyVCS | None:

@X.setter
def X(self, value: Any) -> None:
if value is not None and not isinstance(value, _VCS_TYPES):
if value is not None and not isinstance(value, _X_TYPES):
if sp.issparse(value) or isinstance(value, np.ndarray):
vcls = VCSCArray if isinstance(value, sp.csc_array | sp.csc_matrix) else VCSRArray
value = vcls.from_scipy(value)
else:
_check_vcs_type(value, "X")
_check_vcs_type(value, "X", _X_TYPES)
if (
value is not None
and hasattr(self, "_obs")
Expand Down
182 changes: 182 additions & 0 deletions src/vsparse/_norm_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -608,6 +608,150 @@ def _compute_row_scale(arr: Any, recipe: Recipe) -> np.ndarray:
return row_scale


#: Nonzeros per row block streamed to the device. 2e8 is ~1.6 GB as float32
#: values plus int32 column indices, leaving room on a 24 GB card for the
#: dense operands and temporaries.
DEFAULT_CHUNK_NNZ = 200_000_000


class _DeviceNormalizedView:
"""A normalized view streamed through GPU memory in row blocks.

Supports ``@`` and ``r@`` only, which is the whole contract
``parafac2.backend.GPUMatrix`` asks of a duck-typed matrix. Nothing is
uploaded up front: each product walks the view a block of roughly
``chunk_nnz`` nonzeros at a time, so device residency is bounded by the
block rather than by the dataset. Centering stays the rank-1 correction
:attr:`NormalizedViewBase.means` documents, applied per block.
"""

# NumPy/CuPy would otherwise try to broadcast this object elementwise on
# `lhs @ view` instead of deferring to `__rmatmul__`. `parafac2`'s own
# `GPUMatrix` sets this for the same reason.
__array_priority__ = 1000

__slots__ = (
"_blocks",
"_cache",
"_cache_host",
"_chunk_nnz",
"_means",
"_view",
"dtype",
"shape",
)

def __init__(
self,
view: Any,
chunk_nnz: int = DEFAULT_CHUNK_NNZ,
cache_host: bool = True,
) -> None:
self._view = view
self._means = np.asarray(view.means, dtype=np.float64)
self.shape = view.shape
self.dtype = np.dtype(np.float64)
self._chunk_nnz = chunk_nnz
self._blocks = self._plan_blocks()
# Decoding off the packed view is the expensive part, and a
# compression makes several raw-data passes, so cache the host blocks
# and pay it once.
self._cache_host = cache_host
self._cache: dict[tuple[int, int], Any] = {}

def _plan_blocks(self) -> list[tuple[int, int]]:
"""Contiguous row ranges of roughly ``chunk_nnz`` nonzeros each."""
n_rows = self.shape[0]
if n_rows == 0:
return []
nnz = int(getattr(self._view, "nnz", self._view._arr.nnz))
per_row = max(1.0, nnz / n_rows)
rows = max(1, min(n_rows, int(self._chunk_nnz / per_row)))
return [(s, min(s + rows, n_rows)) for s in range(0, n_rows, rows)]

def _host_block(self, start: int, stop: int) -> Any:
"""One row block as a host CSR with int32 indices, decoded once."""
key = (start, stop)
cached = self._cache.get(key)
if cached is not None:
return cached
host = self._view.select(
slice(start, stop), slice(None), recalculate=False
).to_scipy_sparse(dtype=np.float32)
host = host.tocsr() if hasattr(host, "tocsr") else host
# Under 2**31 nonzeros by construction, so the indices stay int32
# even where the whole matrix would force scipy/CuPy to int64 and
# double what the column indices cost.
host.indices = host.indices.astype(np.int32, copy=False)
host.indptr = host.indptr.astype(np.int32, copy=False)
if self._cache_host:
self._cache[key] = host
return host

def _device_block(self, start: int, stop: int) -> Any:
"""One row block, on device, with int32 indices and canonical flag."""
import cupy as cp # ty: ignore[unresolved-import]
import cupyx.scipy.sparse as cusp # ty: ignore[unresolved-import]

host = self._host_block(start, stop)
block = cusp.csr_matrix(
(cp.asarray(host.data), cp.asarray(host.indices), cp.asarray(host.indptr)),
shape=host.shape,
)
# cuSPARSE rejects a non-canonical CSR rather than canonicalizing one,
# and a matrix rebuilt from raw index arrays carries no canonical flag.
# The block is canonical by construction, so this is a cheap
# device-side check rather than a COO round-trip.
block.has_canonical_format = True
return block

def __matmul__(self, rhs: Any) -> np.ndarray:
"""``self @ rhs``, streamed over row blocks, as a NumPy array.

Host-side because ``parafac2``'s callers do ``np.asarray`` on the
result, which raises on a CuPy array.
"""
import cupy as cp # ty: ignore[unresolved-import]

rhs_arr = np.asarray(rhs)
rhs_1d = rhs_arr.ndim == 1
rhs_2d = rhs_arr[:, None] if rhs_1d else rhs_arr
rhs_d = cp.asarray(rhs_2d, dtype=cp.float32)
shift = cp.asarray(self._means, dtype=cp.float64) @ cp.asarray(rhs_2d, dtype=cp.float64)
out = np.empty((self.shape[0], rhs_2d.shape[1]), dtype=np.float64)
for start, stop in self._blocks:
block = self._device_block(start, stop)
product = cp.asarray(block @ rhs_d, dtype=cp.float64) - shift
out[start:stop] = cp.asnumpy(product)
del block, product
return out.ravel() if rhs_1d else out

def __rmatmul__(self, lhs: Any) -> np.ndarray:
"""``lhs @ self``, streamed over row blocks, as a NumPy array."""
import cupy as cp # ty: ignore[unresolved-import]
import cupyx.cusparse # ty: ignore[unresolved-import]

lhs_arr = np.asarray(lhs)
lhs_1d = lhs_arr.ndim == 1
lhs_2d = lhs_arr[None, :] if lhs_1d else lhs_arr
width = lhs_2d.shape[0]
total = cp.zeros((self.shape[1], width), dtype=cp.float64)
column_weight = cp.zeros(width, dtype=cp.float64)
for start, stop in self._blocks:
block = self._device_block(start, stop)
# `spmm` with the transpose flag, not `dense @ sparse`: CuPy routes
# the latter through `sum_duplicates`, round-tripping the block
# through COO and allocating several times its size. It wants an
# F-contiguous operand.
left = cp.asfortranarray(cp.asarray(lhs_2d[:, start:stop].T, dtype=cp.float32))
total += cp.asarray(cupyx.cusparse.spmm(block, left, transa=True), dtype=cp.float64)
column_weight += cp.asarray(left, dtype=cp.float64).sum(axis=0)
del block, left
total -= cp.outer(cp.asarray(self._means, dtype=cp.float64), column_weight)
out = np.ascontiguousarray(cp.asnumpy(total).T)
return out.ravel() if lhs_1d else out


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

Expand Down Expand Up @@ -873,6 +1017,44 @@ def toarray(self) -> np.ndarray:
)
return out

def copy(self) -> NormalizedViewBase:
"""An independent view over the same base array.

The statistics are copied; the base array is **shared**, because a view
never mutates it and duplicating it would defeat the point of being
lazy -- 15+ GB on a cohort-scale dataset. This mirrors :meth:`select`,
which likewise returns a view sharing the base.

Present so that containers holding a view (see
:class:`~vsparse.VCSCAnnData`) can implement ``copy``/``to_memory``
without materializing. ``anndata`` calls ``to_memory()`` defensively to
realize disk-backed data, which for an in-memory view is a no-op.
"""
return type(self).from_stats(
self._arr,
self.recipe,
np.array(self.a),
np.array(self.b),
np.array(self.c),
np.array(self.s),
stale=self.stale,
)

def to_device(self, backend: str) -> Any:
"""This view, resident on ``backend``'s device.

``backend`` is ``"cpu"``, which returns ``self``, or ``"cuda"``. Only
the sparse ``Delta`` term and the length-``n_cols`` centering vector
move -- see :class:`_DeviceNormalizedView`.
"""
if backend == "cpu":
return self
if backend != "cupy":
raise ValueError(
f"{type(self).__name__}.to_device supports 'cupy' and 'cpu', got {backend!r}."
)
return _DeviceNormalizedView(self)

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

Expand Down
Loading