Skip to content

API Reference

Factorization

scrise.factorization

pf2

pf2(
    X: AnnData | None = None,
    rank: int | None = None,
    random_state=1,
    doEmbedding: bool = True,
    tolerance=1e-06,
    max_iter: int = 100,
    normalize_slices: bool = False,
    backend: str | None = None,
    compress: int
    | tuple[int, int | None]
    | str
    | bool
    | None = None,
    compression_kwarg: dict[str, Any] | None = None,
    parafac2_kwarg: dict[str, Any] | None = None,
    condition_key: str | None = None,
    adata: AnnData | None = None,
) -> anndata.AnnData

Perform PARAFAC2 tensor decomposition on single-cell RNA-seq data.

This is the main function for running RISE analysis. It decomposes the multi-condition single-cell data into condition factors, eigen-state factors, and gene factors, revealing patterns across experimental conditions.

Parameters:

Name Type Description Default
X AnnData

Preprocessed AnnData object containing single-cell RNA-seq data. Must have X.obs["condition_unique_idxs"] indicating which condition each cell belongs to (0-indexed).

None
rank int

Number of components to extract. Determines the complexity of the decomposition. Typically chosen based on variance explained and Factor Match Score analysis (see plot_r2x and plot_fms_diff_ranks).

None
random_state int, optional (default: 1)

Random seed for reproducibility of the decomposition.

1
doEmbedding bool, optional (default: True)

If True, automatically computes PaCMAP embedding of cell projections and stores in X.obsm["X_pf2_PaCMAP"]. This enables visualization functions like plot_labels_pacmap.

True
tolerance float, optional (default: 1e-6)

Convergence threshold for the optimization algorithm. Lower values increase precision but may require more iterations.

1e-06
max_iter int, optional (default: 100)

Maximum number of iterations for the optimization algorithm.

100
normalize_slices bool, optional (default: False)

If True, weights each condition by the inverse of its Frobenius norm, so conditions count equally regardless of cell count. Factors and R2X stay in data units, comparable to an unweighted fit.

False
backend str | None, optional (default: None)

Compute backend to run matrix products on: one of 'mlx', 'cupy', or 'cpu'. If None, the first available accelerator is auto-detected (see :func:~parafac2.backend.get_backend).

None
compress int | tuple[int, int | None] | str | bool | None, optional (default: None)

CANDELINC compression mode passed to parafac2_nd. If None/False (default), exact ALS is used. If "auto" or True, compression dimensions are set automatically from rank. See :func:parafac2.parafac2.parafac2_nd for details.

None
compression_kwarg dict

Additional keyword arguments forwarded to :func:parafac2.compress.compress_dataset (e.g. n_power_iter). Requires compress to also be set.

None
parafac2_kwarg dict

Additional keyword arguments forwarded to parafac2_nd (e.g. n_inner, callback), for options not otherwise exposed here.

None
condition_key str, optional (default: None)

Column in X.obs holding the condition labels, used to derive condition_unique_idxs when that column is not already present.

None
adata anndata.AnnData, optional (default: None)

Alias for X; supply either one, not both.

None

Returns:

Type Description
AnnData

The input AnnData object with added RISE decomposition results:

  • X.uns["Pf2_weights"]: Component weights (shape: rank,)
  • X.uns["Pf2_A"]: Condition factors (shape: n_conditions, rank)
  • X.uns["Pf2_B"]: Eigen-state factors (shape: rank, rank)
  • X.varm["Pf2_C"]: Gene factors (shape: n_genes, rank)
  • X.obsm["projections"]: Cell projections (shape: n_cells, rank)
  • X.obsm["weighted_projections"]: Weighted cell projections (shape: n_cells, rank)
  • X.obsm["X_pf2_PaCMAP"]: PaCMAP embedding (shape: n_cells, 2) if doEmbedding=True

Components are ordered from highest to lowest intrinsic energy (see :func:order_components_by_energy), so that low-energy components -- typically the ones added when the rank is increased -- land at the high end of the ordering.

Source code in scrise/factorization.py
def pf2(
    X: anndata.AnnData | None = None,
    rank: int | None = None,
    random_state=1,
    doEmbedding: bool = True,
    tolerance=1e-6,
    max_iter: int = 100,
    normalize_slices: bool = False,
    backend: str | None = None,
    compress: int | tuple[int, int | None] | str | bool | None = None,
    compression_kwarg: dict[str, Any] | None = None,
    parafac2_kwarg: dict[str, Any] | None = None,
    condition_key: str | None = None,
    adata: anndata.AnnData | None = None,
) -> anndata.AnnData:
    """Perform PARAFAC2 tensor decomposition on single-cell RNA-seq data.

    This is the main function for running RISE analysis. It decomposes the
    multi-condition single-cell data into condition factors, eigen-state factors,
    and gene factors, revealing patterns across experimental conditions.

    Parameters
    ----------
    X : anndata.AnnData
        Preprocessed AnnData object containing single-cell RNA-seq data.
        Must have X.obs["condition_unique_idxs"] indicating which condition
        each cell belongs to (0-indexed).
    rank : int
        Number of components to extract. Determines the complexity of the
        decomposition. Typically chosen based on variance explained and
        Factor Match Score analysis (see plot_r2x and plot_fms_diff_ranks).
    random_state : int, optional (default: 1)
        Random seed for reproducibility of the decomposition.
    doEmbedding : bool, optional (default: True)
        If True, automatically computes PaCMAP embedding of cell projections
        and stores in X.obsm["X_pf2_PaCMAP"]. This enables visualization
        functions like plot_labels_pacmap.
    tolerance : float, optional (default: 1e-6)
        Convergence threshold for the optimization algorithm. Lower values
        increase precision but may require more iterations.
    max_iter : int, optional (default: 100)
        Maximum number of iterations for the optimization algorithm.
    normalize_slices : bool, optional (default: False)
        If True, weights each condition by the inverse of its Frobenius norm,
        so conditions count equally regardless of cell count. Factors and R2X
        stay in data units, comparable to an unweighted fit.
    backend : str | None, optional (default: None)
        Compute backend to run matrix products on: one of ``'mlx'``, ``'cupy'``,
        or ``'cpu'``. If None, the first available accelerator is auto-detected
        (see :func:`~parafac2.backend.get_backend`).
    compress : int | tuple[int, int | None] | str | bool | None, optional (default: None)
        CANDELINC compression mode passed to ``parafac2_nd``. If None/False
        (default), exact ALS is used. If ``"auto"`` or True, compression
        dimensions are set automatically from ``rank``. See
        :func:`parafac2.parafac2.parafac2_nd` for details.
    compression_kwarg : dict, optional
        Additional keyword arguments forwarded to
        :func:`parafac2.compress.compress_dataset` (e.g. ``n_power_iter``).
        Requires ``compress`` to also be set.
    parafac2_kwarg : dict, optional
        Additional keyword arguments forwarded to ``parafac2_nd`` (e.g.
        ``n_inner``, ``callback``), for options not otherwise exposed here.

    condition_key : str, optional (default: None)
        Column in ``X.obs`` holding the condition labels, used to derive
        ``condition_unique_idxs`` when that column is not already present.
    adata : anndata.AnnData, optional (default: None)
        Alias for ``X``; supply either one, not both.

    Returns
    -------
    anndata.AnnData
        The input AnnData object with added RISE decomposition results:

        - X.uns["Pf2_weights"]: Component weights (shape: rank,)
        - X.uns["Pf2_A"]: Condition factors (shape: n_conditions, rank)
        - X.uns["Pf2_B"]: Eigen-state factors (shape: rank, rank)
        - X.varm["Pf2_C"]: Gene factors (shape: n_genes, rank)
        - X.obsm["projections"]: Cell projections (shape: n_cells, rank)
        - X.obsm["weighted_projections"]: Weighted cell projections
          (shape: n_cells, rank)
        - X.obsm["X_pf2_PaCMAP"]: PaCMAP embedding (shape: n_cells, 2)
          if doEmbedding=True

        Components are ordered from highest to lowest intrinsic energy
        (see :func:`order_components_by_energy`), so that low-energy
        components -- typically the ones added when the rank is increased
        -- land at the high end of the ordering.
    """
    if X is None and adata is not None:
        X = adata
    if X is None:
        raise ValueError("Either X or adata must be provided.")
    if rank is None:
        raise ValueError("rank must be provided.")

    if "condition_unique_idxs" not in X.obs:
        if condition_key is not None and condition_key in X.obs:
            X.obs["condition_unique_idxs"] = pd.Categorical(X.obs[condition_key]).codes
        else:
            raise KeyError(
                "X.obs must contain 'condition_unique_idxs', or provide 'condition_key' pointing to a valid column in X.obs."
            )

    pf_out, _ = run_parafac2(
        X,
        rank=rank,
        random_state=random_state,
        tol=tolerance,
        n_iter_max=max_iter,
        normalize_slices=normalize_slices,
        backend=backend,
        compress=compress,
        compression_kwarg=compression_kwarg,
        parafac2_kwarg=parafac2_kwarg,
    )

    X = store_pf2(X, pf_out)
    X = order_components_by_energy(X)

    if doEmbedding:
        pcm = PaCMAP(random_state=random_state)
        X.obsm["X_pf2_PaCMAP"] = pcm.fit_transform(X.obsm["projections"])

    return X

correct_conditions

correct_conditions(X: AnnData)

Correct the condition factors by normalizing for overall read depth.

This function adjusts condition factors (stored in X.uns["Pf2_A"]) to account for differences in sequencing depth across conditions. It uses linear regression to model the relationship between total read counts and condition factor magnitudes, then applies a correction.

Parameters:

Name Type Description Default
X AnnData

AnnData object containing RISE decomposition results. Must have: - X.obs["condition_unique_idxs"]: 0-indexed condition assignments - X.uns["Pf2_A"]: Condition factors from PARAFAC2 decomposition

required

Returns:

Type Description
ndarray

Corrected condition factors normalized by sequencing depth

Source code in scrise/factorization.py
def correct_conditions(X: anndata.AnnData):
    """Correct the condition factors by normalizing for overall read depth.

    This function adjusts condition factors (stored in X.uns["Pf2_A"]) to account for
    differences in sequencing depth across conditions. It uses linear regression to
    model the relationship between total read counts and condition factor magnitudes,
    then applies a correction.

    Parameters
    ----------
    X : anndata.AnnData
        AnnData object containing RISE decomposition results. Must have:
        - X.obs["condition_unique_idxs"]: 0-indexed condition assignments
        - X.uns["Pf2_A"]: Condition factors from PARAFAC2 decomposition

    Returns
    -------
    numpy.ndarray
        Corrected condition factors normalized by sequencing depth
    """
    sgIndex = np.asarray(X.obs["condition_unique_idxs"])
    cond_mean = gmean(X.uns["Pf2_A"], axis=1)

    if X.X is None:
        raise TypeError("X.X must not be None.")
    # X.X's declared type is a large union of array-like/backed-storage types
    # from the AnnData stub; at runtime this is always a dense or sparse
    # in-memory array supporting `.sum`.
    x_count = np.asarray(cast(Any, X.X).sum(axis=1)).ravel()

    n_conds = int(np.amax(sgIndex)) + 1
    counts = np.bincount(sgIndex, weights=x_count, minlength=n_conds).reshape(-1, 1)

    lr = LinearRegression()
    lr.fit(counts, cond_mean.reshape(-1, 1))

    counts_correct = lr.predict(counts)

    return X.uns["Pf2_A"] / counts_correct

order_components_by_energy

order_components_by_energy(X: AnnData) -> anndata.AnnData

Reorder PARAFAC2 components by intrinsic energy.

RISE previously inherited an ordering of components based on the Gini coefficient (variance-to-mean ratio) of the condition factor. That ordering is a heuristic: it is not particularly stable, and it does not guarantee that, when moving from rank N to rank N+1, the first N components of the new fit correspond to the N components of the old fit.

This function instead orders components by their intrinsic energy, |weights[r]| * ||A[:, r]|| * ||C[:, r]|| (the product of component weights, condition-factor, and gene-factor column norms), which is directly determined by the fit and is not subject to arbitrary rescaling. Components are ordered from highest to lowest energy, so that low-energy components -- which tend to be the ones added when the rank is increased -- land at the high end of the ordering. This is consistent with the expectation, motivated by Harshman's Uniqueness Theorem for PARAFAC2 (which fixes the decomposition up to permutation and sign/scale given >= 3 conditions and full-column-rank factors), that a rank-(N+1) fit should mostly agree with a rank-N fit on its first N components.

Before computing the ordering, a canonical sign is imposed on each component (see :func:canonical_component_signs) so that the ordering, and any downstream comparison across fits, is well defined.

Parameters:

Name Type Description Default
X AnnData

AnnData object containing RISE decomposition results (as produced by :func:pf2). Must contain X.uns["Pf2_A"], X.uns["Pf2_B"], X.uns["Pf2_weights"], and X.varm["Pf2_C"].

required

Returns:

Type Description
AnnData

The same AnnData object, with Pf2_A, Pf2_B, Pf2_C, Pf2_weights (and, if present, obsm["projections"] and obsm["weighted_projections"]) reordered/updated in place. The eigen-state axis of Pf2_B, and the matching columns of obsm["projections"], are permuted alongside the components so that the maximal-diagonal form of B established by parafac2.utils.standardize_pf2 is preserved.

Source code in scrise/factorization.py
def order_components_by_energy(X: anndata.AnnData) -> anndata.AnnData:
    """Reorder PARAFAC2 components by intrinsic energy.

    RISE previously inherited an ordering of components based on the Gini
    coefficient (variance-to-mean ratio) of the condition factor. That
    ordering is a heuristic: it is not particularly stable, and it does not
    guarantee that, when moving from rank N to rank N+1, the first N
    components of the new fit correspond to the N components of the old fit.

    This function instead orders components by their intrinsic energy,
    ``|weights[r]| * ||A[:, r]|| * ||C[:, r]||`` (the product of component
    weights, condition-factor, and gene-factor column norms), which is
    directly determined by the fit and is not subject to arbitrary rescaling.
    Components are ordered from highest to lowest energy, so that low-energy
    components -- which tend to be the ones added when the rank is
    increased -- land at the high end of the ordering. This is consistent
    with the expectation, motivated by Harshman's Uniqueness Theorem for
    PARAFAC2 (which fixes the decomposition up to permutation and
    sign/scale given >= 3 conditions and full-column-rank factors), that a
    rank-(N+1) fit should mostly agree with a rank-N fit on its first N
    components.

    Before computing the ordering, a canonical sign is imposed on each
    component (see :func:`canonical_component_signs`) so that the ordering,
    and any downstream comparison across fits, is well defined.

    Parameters
    ----------
    X : anndata.AnnData
        AnnData object containing RISE decomposition results (as produced by
        :func:`pf2`). Must contain X.uns["Pf2_A"], X.uns["Pf2_B"],
        X.uns["Pf2_weights"], and X.varm["Pf2_C"].

    Returns
    -------
    anndata.AnnData
        The same AnnData object, with Pf2_A, Pf2_B, Pf2_C, Pf2_weights
        (and, if present, obsm["projections"] and
        obsm["weighted_projections"]) reordered/updated in place. The
        eigen-state axis of Pf2_B, and the matching columns of
        obsm["projections"], are permuted alongside the components so that
        the maximal-diagonal form of B established by
        ``parafac2.utils.standardize_pf2`` is preserved.
    """
    A = np.array(X.uns["Pf2_A"])
    B = np.array(X.uns["Pf2_B"])
    C = np.array(X.varm["Pf2_C"])
    weights = np.array(X.uns["Pf2_weights"])

    # Canonical sign convention, shared between the condition factor (A) and
    # the gene factor (C); the eigen-state factor (B) is left as the
    # unflipped reference so that the product A x B x C is unchanged.
    signs = canonical_component_signs(C)
    A = A * signs
    C = C * signs

    energy = np.abs(weights) * np.linalg.norm(A, axis=0) * np.linalg.norm(C, axis=0)
    order = np.argsort(energy)[::-1]

    X.uns["Pf2_A"] = A[:, order]
    X.varm["Pf2_C"] = C[:, order]
    X.uns["Pf2_weights"] = weights[order]

    # B is indexed by (eigen-state, component). ``parafac2.utils.standardize_pf2``
    # permutes B's *rows* so that its diagonal is maximal, pairing eigen-state i
    # with component i, and permutes the columns of each projection to match.
    # Permuting only B's columns here would move each diagonal entry off the
    # diagonal and destroy that pairing, so the eigen-state axis is relabeled by
    # the same permutation. Since P_k B is left unchanged by this relabeling
    # (up to the component permutation), the reconstruction is untouched.
    X.uns["Pf2_B"] = B[np.ix_(order, order)]

    if "projections" in X.obsm:
        X.obsm["projections"] = np.asarray(X.obsm["projections"])[:, order]
        X.obsm["weighted_projections"] = (
            X.obsm["projections"] @ X.uns["Pf2_B"]
        ).astype(np.float32, copy=False)

    return X

canonical_component_signs

canonical_component_signs(C: ndarray) -> np.ndarray

Compute a canonical sign for each component of a gene factor matrix.

PARAFAC2 components are only identified up to a sign flip that is shared between two of the three factor matrices (Harshman's Uniqueness Theorem fixes the decomposition up to permutation and sign/scale). This makes any comparison across fits (different ranks, different random seeds, etc.) ambiguous unless a canonical sign is first imposed. We adopt the convention that the largest-magnitude entry of each gene-factor (C) column should be positive.

Parameters:

Name Type Description Default
C ndarray

Gene factor matrix, shape (n_genes, rank).

required

Returns:

Type Description
ndarray

Array of +1/-1 signs, shape (rank,), one per component.

Source code in scrise/factorization.py
def canonical_component_signs(C: np.ndarray) -> np.ndarray:
    """Compute a canonical sign for each component of a gene factor matrix.

    PARAFAC2 components are only identified up to a sign flip that is shared
    between two of the three factor matrices (Harshman's Uniqueness Theorem
    fixes the decomposition up to permutation and sign/scale). This makes any
    comparison across fits (different ranks, different random seeds, etc.)
    ambiguous unless a canonical sign is first imposed. We adopt the
    convention that the largest-magnitude entry of each gene-factor (``C``)
    column should be positive.

    Parameters
    ----------
    C : numpy.ndarray
        Gene factor matrix, shape (n_genes, rank).

    Returns
    -------
    numpy.ndarray
        Array of +1/-1 signs, shape (rank,), one per component.
    """
    rank = C.shape[1]
    max_idx = np.argmax(np.abs(C), axis=0)
    signs = np.ones(rank)
    signs[C[max_idx, np.arange(rank)] < 0] = -1.0
    return signs

match_components_across_ranks

match_components_across_ranks(
    C_low: ndarray, C_high: ndarray, threshold: float = 0.6
) -> tuple[np.ndarray, np.ndarray]

Match components between two PARAFAC2 fits of adjacent rank by cosine similarity of their gene factors.

This provides the "does component X at rank N correspond to component Y at rank N-1" matching primitive described in issue #520. It is a small, self-contained addition, not a full cross-rank benchmarking pipeline: given the gene factors of a rank-N fit and a rank-(N+1) fit (each already sign-canonicalized, e.g. via :func:canonical_component_signs), it performs Hungarian maximum-weight matching on cosine similarity and reports which components matched, and which rank-(N+1) component(s) had no good match (i.e. are candidates for the newly added component).

Parameters:

Name Type Description Default
C_low ndarray

Gene factor matrix of the lower-rank fit, shape (n_genes, rank_low).

required
C_high ndarray

Gene factor matrix of the higher-rank fit, shape (n_genes, rank_high), with rank_high >= rank_low.

required
threshold float, optional (default: 0.6)

Minimum cosine similarity for two components to be considered a match.

0.6

Returns:

Type Description
tuple of numpy.ndarray

matched_pairs : numpy.ndarray of shape (n_matches, 2) Each row is (index in C_low, index in C_high) for a matched pair of components. unmatched_high : numpy.ndarray Indices into C_high of components with no match above threshold -- candidates for the newly added component(s).

Source code in scrise/factorization.py
def match_components_across_ranks(
    C_low: np.ndarray, C_high: np.ndarray, threshold: float = 0.6
) -> tuple[np.ndarray, np.ndarray]:
    """Match components between two PARAFAC2 fits of adjacent rank by cosine
    similarity of their gene factors.

    This provides the "does component X at rank N correspond to component Y
    at rank N-1" matching primitive described in issue #520. It is a small,
    self-contained addition, not a full cross-rank benchmarking pipeline:
    given the gene factors of a rank-N fit and a rank-(N+1) fit (each
    already sign-canonicalized, e.g. via :func:`canonical_component_signs`),
    it performs Hungarian maximum-weight matching on cosine similarity and
    reports which components matched, and which rank-(N+1) component(s) had
    no good match (i.e. are candidates for the newly added component).

    Parameters
    ----------
    C_low : numpy.ndarray
        Gene factor matrix of the lower-rank fit, shape (n_genes, rank_low).
    C_high : numpy.ndarray
        Gene factor matrix of the higher-rank fit, shape (n_genes,
        rank_high), with rank_high >= rank_low.
    threshold : float, optional (default: 0.6)
        Minimum cosine similarity for two components to be considered a
        match.

    Returns
    -------
    tuple of numpy.ndarray
        matched_pairs : numpy.ndarray of shape (n_matches, 2)
            Each row is (index in C_low, index in C_high) for a matched
            pair of components.
        unmatched_high : numpy.ndarray
            Indices into C_high of components with no match above
            ``threshold`` -- candidates for the newly added component(s).
    """
    from scipy.optimize import linear_sum_assignment

    C_low_n = C_low / np.linalg.norm(C_low, axis=0, keepdims=True)
    C_high_n = C_high / np.linalg.norm(C_high, axis=0, keepdims=True)

    cos_sim = C_low_n.T @ C_high_n  # (rank_low, rank_high)

    row_ind, col_ind = linear_sum_assignment(-cos_sim)

    matched_mask = cos_sim[row_ind, col_ind] >= threshold
    matched_pairs = np.stack([row_ind[matched_mask], col_ind[matched_mask]], axis=1)

    unmatched_high = np.setdiff1d(
        np.arange(C_high.shape[1]), matched_pairs[:, 1] if matched_pairs.size else []
    )

    return matched_pairs, unmatched_high

Factor Import/Export

scrise.factor_io

Reading and writing RISE factors, without the raw expression matrix.

Factors are exported to h5ad with the projection matrix OPQ-quantized and the cell barcodes packed into a uint8 matrix, both of which have to be undone on load.

export_factors

export_factors(
    X: AnnData,
    filename: str,
    fidelity_threshold: float = 0.99,
    random_state: int = 42,
) -> anndata.AnnData

Export RISE decomposition factors to an h5ad file without raw expression data.

Compresses the projection matrix using Optimized Product Quantization (OPQ) to meet or exceed the specified fidelity threshold (R^2 >= fidelity_threshold). All factor matrices (Pf2_A, Pf2_B, Pf2_weights, Pf2_C) are stored in float32. Weighted projections are never stored on disk because they can be reconstructed deterministically as projections @ Pf2_B. PaCMAP embeddings are optionally stored if present.

Parameters:

Name Type Description Default
X AnnData

AnnData object containing RISE decomposition results. Must contain: - X.uns["Pf2_A"], X.uns["Pf2_B"], X.uns["Pf2_weights"] - X.varm["Pf2_C"] - X.obsm["projections"]

required
filename str

Output file path (.h5ad).

required
fidelity_threshold float, optional (default: 0.99)

Target R^2 reconstruction accuracy threshold for projection compression.

0.99
random_state int, optional (default: 42)

Random seed for reproducibility during OPQ codebook training.

42

Returns:

Type Description
AnnData

The factor-only AnnData object written to disk.

Source code in scrise/factor_io.py
def export_factors(
    X: anndata.AnnData,
    filename: str,
    fidelity_threshold: float = 0.99,
    random_state: int = 42,
) -> anndata.AnnData:
    """Export RISE decomposition factors to an h5ad file without raw expression data.

    Compresses the projection matrix using Optimized Product Quantization (OPQ)
    to meet or exceed the specified fidelity threshold (R^2 >= fidelity_threshold).
    All factor matrices (Pf2_A, Pf2_B, Pf2_weights, Pf2_C) are stored in float32.
    Weighted projections are never stored on disk because they can be reconstructed
    deterministically as projections @ Pf2_B.
    PaCMAP embeddings are optionally stored if present.

    Parameters
    ----------
    X : anndata.AnnData
        AnnData object containing RISE decomposition results. Must contain:
        - X.uns["Pf2_A"], X.uns["Pf2_B"], X.uns["Pf2_weights"]
        - X.varm["Pf2_C"]
        - X.obsm["projections"]
    filename : str
        Output file path (.h5ad).
    fidelity_threshold : float, optional (default: 0.99)
        Target R^2 reconstruction accuracy threshold for projection compression.
    random_state : int, optional (default: 42)
        Random seed for reproducibility during OPQ codebook training.

    Returns
    -------
    anndata.AnnData
        The factor-only AnnData object written to disk.
    """
    _require_export_factors(X)

    # Factor matrices in float32
    uns_dict = _floats_to_float32(X.uns)
    uns_dict["Pf2_A"] = np.asarray(X.uns["Pf2_A"], dtype=np.float32)
    uns_dict["Pf2_B"] = np.asarray(X.uns["Pf2_B"], dtype=np.float32)
    uns_dict["Pf2_weights"] = np.asarray(X.uns["Pf2_weights"], dtype=np.float32)

    varm_dict = _floats_to_float32(X.varm)
    varm_dict["Pf2_C"] = np.asarray(X.varm["Pf2_C"], dtype=np.float32)

    codes, opq_uns = _compress_projections(X, fidelity_threshold, random_state)
    uns_dict.update(opq_uns)

    # obsm excludes weighted_projections and the uncompressed projections.
    obsm_dict = {"projections_opq_codes": codes.astype(np.uint8)}
    obsm_dict.update(_embedding_for_export(X))

    obs_df, var_df = X.obs, X.var
    if not isinstance(obs_df, pd.DataFrame) or not isinstance(var_df, pd.DataFrame):
        raise TypeError(
            "X.obs and X.var must be in-memory pandas DataFrames "
            "(backed Dataset2D is not supported)."
        )
    obs = obs_df.copy()
    n_cells = len(obs)
    packed_names = _pack_obs_names(obs)
    if packed_names is not None:
        uns_dict["_obs_names_bytes"] = packed_names

    factors_adata = anndata.AnnData(
        obs=obs,
        var=var_df.copy(),
        uns=uns_dict,
        varm=cast(Mapping[str, Sequence[Any]], varm_dict),
        obsm=cast(Mapping[str, Sequence[Any]], obsm_dict),
    )

    out_dir = os.path.dirname(os.path.abspath(filename))
    if out_dir:
        os.makedirs(out_dir, exist_ok=True)
    factors_adata.write_h5ad(filename)

    if "_obs_names_bytes" in uns_dict:
        _recompress_obs_names(filename, n_cells)
    return factors_adata

load_factors

load_factors(
    filename: str, raw_path: str | None = None
) -> anndata.AnnData

Load RISE decomposition factors from an h5ad file, decompressing OPQ projections and optionally rebuilding the full dataset from raw expression data.

Parameters:

Name Type Description Default
filename str

Path to factors .h5ad file.

required
raw_path str

Path to raw AnnData or IVCSR .h5ad / .h5 file. If provided, cells and genes are matched based on cell barcodes and gene names to attach the expression matrix.

None

Returns:

Type Description
AnnData

AnnData object with reconstructed projections, weighted_projections, factors, and optionally the raw expression matrix X.

Source code in scrise/factor_io.py
def load_factors(
    filename: str,
    raw_path: str | None = None,
) -> anndata.AnnData:
    """Load RISE decomposition factors from an h5ad file, decompressing OPQ projections
    and optionally rebuilding the full dataset from raw expression data.

    Parameters
    ----------
    filename : str
        Path to factors .h5ad file.
    raw_path : str, optional
        Path to raw AnnData or IVCSR .h5ad / .h5 file. If provided, cells and genes
        are matched based on cell barcodes and gene names to attach the expression matrix.

    Returns
    -------
    anndata.AnnData
        AnnData object with reconstructed projections, weighted_projections,
        factors, and optionally the raw expression matrix X.
    """
    adata = anndata.read_h5ad(filename)
    _restore_obs_names(adata)
    _restore_projections(adata)

    if raw_path is not None:
        _attach_raw_data(adata, raw_path)

    return adata

Rank Selection

scrise.rank_selection

Rank selection for RISE via bi-cross-validation (BiCV).

Bi-cross-validation extends ordinary cross-validation to two-way (row and column) held-out blocks. For RISE, we hold out a random subset of cells and a random subset of genes, fit PARAFAC2 on the remaining (train-cell x train-gene) block, and then measure how well the fitted model predicts the held-out (test-cell x test-gene) block. Unlike the ordinary in-sample fit R2X (which increases monotonically with rank), the BiCV R2X penalizes overfitting and typically peaks near the "true" rank of the data.

bicv

bicv(
    X: AnnData | None = None,
    ranks: Sequence[int] | None = None,
    n_repeats: int = 3,
    held_out_cell_frac: float = 0.5,
    held_out_gene_frac: float = 0.5,
    random_state: int | None = None,
    tolerance: float = 1e-06,
    max_iter: int = 200,
    compress: int
    | tuple[int, int | None]
    | str
    | bool
    | None = "auto",
    compression_kwarg: dict[str, Any] | None = None,
    parafac2_kwarg: dict[str, Any] | None = None,
    condition_key: str | None = None,
    adata: AnnData | None = None,
) -> pd.DataFrame

Evaluate rank via bi-cross-validation (BiCV) and in-sample fit R2X.

For each candidate rank, computes both the ordinary in-sample fit R2X (using the full dataset, as in :func:scrise.factorization.rise_pca_r2x) and the BiCV R2X (n_repeats independent random cell/gene splits, each evaluated at every rank -- see :func:_bicv_trial). The fit R2X increases monotonically with rank; the BiCV R2X penalizes overfitting and typically peaks near the rank that best generalizes to held-out data. Plot both with :func:scrise.plotting.plot_bicv_r2x to select a rank.

Each of the n_repeats splits (and the full dataset, for the in-sample fit) is compressed once -- sized for the largest rank requested -- and every rank's PARAFAC2 fit reuses that same compression, rather than recompressing per rank (see :func:_fit_at_ranks). Compression is by far the most expensive, O(nnz) step against the raw data; fitting a smaller rank from an already-compressed representation only touches the small dense compressed cores. This applies whenever compression_kwarg is given (as it must be to control e.g. n_power_iter); without it, each rank still goes through its own call into parafac2_nd's internal compression shortcut, unchanged.

Parameters:

Name Type Description Default
X AnnData

Preprocessed AnnData object containing single-cell RNA-seq data. Must have X.obs["condition_unique_idxs"] and X.var["means"] (as produced by parafac2.normalize.prepare_dataset).

None
ranks sequence of int

Candidate rank values to evaluate (e.g., [5, 10, 15, 20, 25, 30]).

None
n_repeats int, optional (default: 3)

Number of independent random cell/gene splits, each evaluated at every rank in ranks. Higher values give a less noisy BiCV estimate but take longer.

3
held_out_cell_frac float, optional (default: 0.5)

Fraction of cells held out per condition in each BiCV trial.

0.5
held_out_gene_frac float, optional (default: 0.5)

Fraction of genes held out in each BiCV trial.

0.5
random_state int

Random seed for reproducibility.

None
tolerance float, optional (default: 1e-6)

Convergence threshold passed to the PARAFAC2 fit.

1e-06
max_iter int, optional (default: 200)

Maximum number of iterations passed to the PARAFAC2 fit.

200
compress int | tuple[int, int | None] | str | bool | None

CANDELINC compression mode passed to each PARAFAC2 fit. Defaults to "auto" (compression dimensions set per rank), which sharply cuts the cost of sweeping many ranks and repeats over raw data. Pass None/False to fall back to exact ALS.

'auto'
compression_kwarg dict

Additional keyword arguments forwarded to :func:parafac2.compress.compress_dataset (e.g. n_power_iter) for every trial and in-sample fit. Requires compress to also be set.

None
parafac2_kwarg dict

Additional keyword arguments forwarded to parafac2_nd for every trial and in-sample fit (e.g. normalize_slices, backend, n_inner), for underlying PARAFAC2 options not otherwise exposed here. See :func:scrise.factorization.pf2's normalize_slices for how it can help with unequal cell counts across conditions.

None
condition_key str, optional (default: None)

Column in X.obs holding the condition labels, used to derive condition_unique_idxs when that column is not already present.

None
adata anndata.AnnData, optional (default: None)

Alias for X; supply either one, not both.

None

Returns:

Type Description
DataFrame

Long-form DataFrame with columns "Rank", "Repeat", "Metric" (one of "Fit R2X" or "BiCV R2X"), and "R2X". Ready to pass to :func:scrise.plotting.plot_bicv_r2x.

BiCV rows carry per-trial diagnostics as additional columns:

Train Block R2X In-sample R2X on the block the model was actually fit to. NTrainGenes, NTestGenes, NTrainCells, NTestCells The realised block sizes, which set the scale of the spread across repeats. Seed The seed that determines the trial's splits and initialisation.

These columns are NaN on "Fit R2X" rows, which come from an unsplit fit on the full dataset.

Source code in scrise/rank_selection.py
def bicv(
    X: anndata.AnnData | None = None,
    ranks: Sequence[int] | None = None,
    n_repeats: int = 3,
    held_out_cell_frac: float = 0.5,
    held_out_gene_frac: float = 0.5,
    random_state: int | None = None,
    tolerance: float = 1e-6,
    max_iter: int = 200,
    compress: int | tuple[int, int | None] | str | bool | None = "auto",
    compression_kwarg: dict[str, Any] | None = None,
    parafac2_kwarg: dict[str, Any] | None = None,
    condition_key: str | None = None,
    adata: anndata.AnnData | None = None,
) -> pd.DataFrame:
    """Evaluate rank via bi-cross-validation (BiCV) and in-sample fit R2X.

    For each candidate rank, computes both the ordinary in-sample fit R2X
    (using the full dataset, as in :func:`scrise.factorization.rise_pca_r2x`)
    and the BiCV R2X (``n_repeats`` independent random cell/gene splits,
    each evaluated at every rank -- see :func:`_bicv_trial`). The fit R2X
    increases monotonically with rank; the BiCV R2X penalizes overfitting and
    typically peaks near the rank that best generalizes to held-out data.
    Plot both with :func:`scrise.plotting.plot_bicv_r2x` to select a rank.

    Each of the ``n_repeats`` splits (and the full dataset, for the in-sample
    fit) is compressed once -- sized for the *largest* rank requested -- and
    every rank's PARAFAC2 fit reuses that same compression, rather than
    recompressing per rank (see :func:`_fit_at_ranks`). Compression is by far
    the most expensive, ``O(nnz)`` step against the raw data; fitting a
    smaller rank from an already-compressed representation only touches the
    small dense compressed cores. This applies whenever ``compression_kwarg``
    is given (as it must be to control e.g. ``n_power_iter``); without it,
    each rank still goes through its own call into ``parafac2_nd``'s internal
    compression shortcut, unchanged.

    Parameters
    ----------
    X : anndata.AnnData
        Preprocessed AnnData object containing single-cell RNA-seq data.
        Must have X.obs["condition_unique_idxs"] and X.var["means"]
        (as produced by ``parafac2.normalize.prepare_dataset``).
    ranks : sequence of int
        Candidate rank values to evaluate (e.g., [5, 10, 15, 20, 25, 30]).
    n_repeats : int, optional (default: 3)
        Number of independent random cell/gene splits, each evaluated at
        every rank in ``ranks``. Higher values give a less noisy BiCV
        estimate but take longer.
    held_out_cell_frac : float, optional (default: 0.5)
        Fraction of cells held out per condition in each BiCV trial.
    held_out_gene_frac : float, optional (default: 0.5)
        Fraction of genes held out in each BiCV trial.
    random_state : int, optional
        Random seed for reproducibility.
    tolerance : float, optional (default: 1e-6)
        Convergence threshold passed to the PARAFAC2 fit.
    max_iter : int, optional (default: 200)
        Maximum number of iterations passed to the PARAFAC2 fit.
    compress : int | tuple[int, int | None] | str | bool | None, optional
        CANDELINC compression mode passed to each PARAFAC2 fit. Defaults to
        ``"auto"`` (compression dimensions set per rank), which sharply cuts
        the cost of sweeping many ranks and repeats over raw data. Pass
        None/False to fall back to exact ALS.
    compression_kwarg : dict, optional
        Additional keyword arguments forwarded to
        :func:`parafac2.compress.compress_dataset` (e.g. ``n_power_iter``)
        for every trial and in-sample fit. Requires ``compress`` to also be
        set.
    parafac2_kwarg : dict, optional
        Additional keyword arguments forwarded to ``parafac2_nd`` for every
        trial and in-sample fit (e.g. ``normalize_slices``, ``backend``,
        ``n_inner``), for underlying PARAFAC2 options not otherwise exposed
        here. See :func:`scrise.factorization.pf2`'s ``normalize_slices``
        for how it can help with unequal cell counts across conditions.

    condition_key : str, optional (default: None)
        Column in ``X.obs`` holding the condition labels, used to derive
        ``condition_unique_idxs`` when that column is not already present.
    adata : anndata.AnnData, optional (default: None)
        Alias for ``X``; supply either one, not both.

    Returns
    -------
    pandas.DataFrame
        Long-form DataFrame with columns "Rank", "Repeat", "Metric" (one of
        "Fit R2X" or "BiCV R2X"), and "R2X". Ready to pass to
        :func:`scrise.plotting.plot_bicv_r2x`.

        BiCV rows carry per-trial diagnostics as additional columns:

        ``Train Block R2X``
            In-sample R2X on the block the model was actually fit to.
        ``NTrainGenes``, ``NTestGenes``, ``NTrainCells``, ``NTestCells``
            The realised block sizes, which set the scale of the spread across
            repeats.
        ``Seed``
            The seed that determines the trial's splits and initialisation.

        These columns are NaN on "Fit R2X" rows, which come from an
        unsplit fit on the full dataset.
    """
    X, ranks = _resolve_bicv_inputs(
        X,
        adata,
        ranks,
        n_repeats,
        held_out_cell_frac,
        held_out_gene_frac,
        condition_key,
    )

    rng = np.random.default_rng(random_state)
    rows = []

    # The full (unsplit) dataset is the same for every rank, so -- like the
    # per-trial train blocks below -- it's compressed once (sized for the
    # largest rank) and reused, rather than once per rank.
    fit_results = _fit_at_ranks(
        X, ranks, tolerance, max_iter, compress, compression_kwarg, parafac2_kwarg, rng
    )
    for rank, (_, fit_r2x) in zip(ranks, fit_results, strict=True):
        # The full-data fit has no split, so the per-trial columns are absent
        # here rather than zero.
        rows.append({"Rank": rank, "Repeat": 0, "Metric": "Fit R2X", "R2X": fit_r2x})

    # One held-out split per repeat, evaluated at every rank (see
    # `_bicv_trial`), rather than one independent split per (rank, repeat)
    # pair -- the split, and the compression of its train block, no longer
    # need to be redone for each rank.
    for repeat in tqdm(range(n_repeats), desc="BiCV repeats"):
        trial_seed = int(rng.integers(np.iinfo(np.int32).max))
        trial_rows = _bicv_trial(
            X,
            ranks,
            held_out_cell_frac,
            held_out_gene_frac,
            trial_seed,
            tolerance,
            max_iter,
            compress,
            compression_kwarg,
            parafac2_kwarg,
        )
        for trial in trial_rows:
            rank = trial.pop("Rank")
            rows.append(
                {
                    "Rank": rank,
                    "Repeat": repeat,
                    "Metric": "BiCV R2X",
                    "R2X": trial.pop("BiCV R2X"),
                    **trial,
                }
            )

    results = pd.DataFrame(rows)

    bicv_means = results[results["Metric"] == "BiCV R2X"].groupby("Rank")["R2X"].mean()
    best_rank = int(bicv_means.idxmax())
    if best_rank in (ranks[0], ranks[-1]):
        warnings.warn(
            f"bicv: the rank with the highest mean BiCV R2X ({best_rank}) is at "
            f"the edge of the tested ranks ({ranks[0]}-{ranks[-1]}). The true "
            "optimum may lie outside this range -- consider testing additional "
            "ranks beyond it.",
            stacklevel=2,
        )

    return results

Annotation Alignment

scrise.annotation_alignment

Cell-type alignment scoring for RISE component projections.

Quantifies how well a RISE component's cell loadings / projections align with annotated cell types, answering: 1. Uniqueness: Does the component concentrate on a single annotated cell type (tau)? 2. Combination alignment: Does it align with a specific subset of cell types (AUROC + FDR, eta^2)?

CellTypeAlignmentResults dataclass

Alignment results across multiple RISE components.

Parameters:

Name Type Description Default
results list[ComponentAlignmentResult]

List of alignment results for individual components.

required
enrichment DataFrame

Components x cell types matrix of AUROC values.

required
p_values DataFrame

Components x cell types matrix of permutation p-values.

required
q_values DataFrame

Components x cell types matrix of joint BH FDR-adjusted p-values.

required
tau Series

Tissue specificity index tau for each component.

required
eta_squared Series

Eta-squared (variance explained) for each component.

required
kruskal_epsilon_squared Series

Kruskal-Wallis epsilon-squared for each component.

required
significant_cell_types dict[int | str, list[str]]

Mapping of component to significant cell types.

required
alpha float

Significance threshold used for q-values.

0.05
Source code in scrise/annotation_alignment.py
@dataclass
class CellTypeAlignmentResults:
    """Alignment results across multiple RISE components.

    Parameters
    ----------
    results : list[ComponentAlignmentResult]
        List of alignment results for individual components.
    enrichment : pd.DataFrame
        Components x cell types matrix of AUROC values.
    p_values : pd.DataFrame
        Components x cell types matrix of permutation p-values.
    q_values : pd.DataFrame
        Components x cell types matrix of joint BH FDR-adjusted p-values.
    tau : pd.Series
        Tissue specificity index tau for each component.
    eta_squared : pd.Series
        Eta-squared (variance explained) for each component.
    kruskal_epsilon_squared : pd.Series
        Kruskal-Wallis epsilon-squared for each component.
    significant_cell_types : dict[int | str, list[str]]
        Mapping of component to significant cell types.
    alpha : float
        Significance threshold used for q-values.
    """

    results: list[ComponentAlignmentResult]
    enrichment: pd.DataFrame
    p_values: pd.DataFrame
    q_values: pd.DataFrame
    tau: pd.Series
    eta_squared: pd.Series
    kruskal_epsilon_squared: pd.Series
    significant_cell_types: dict[int | str, list[str]]
    alpha: float = 0.05

    def summary(self) -> pd.DataFrame:
        """Return a summary DataFrame across all components."""
        top_types = []
        for comp in self.enrichment.index:
            row = self.enrichment.loc[comp]
            top_types.append(row.idxmax())

        sig_str = [
            ", ".join(self.significant_cell_types.get(comp, []))
            for comp in self.enrichment.index
        ]

        df = pd.DataFrame(
            {
                "tau": self.tau,
                "eta_squared": self.eta_squared,
                "kruskal_epsilon_squared": self.kruskal_epsilon_squared,
                "top_cell_type": top_types,
                "significant_cell_types": sig_str,
            },
            index=self.enrichment.index,
        )
        return df

summary

summary() -> pd.DataFrame

Return a summary DataFrame across all components.

Source code in scrise/annotation_alignment.py
def summary(self) -> pd.DataFrame:
    """Return a summary DataFrame across all components."""
    top_types = []
    for comp in self.enrichment.index:
        row = self.enrichment.loc[comp]
        top_types.append(row.idxmax())

    sig_str = [
        ", ".join(self.significant_cell_types.get(comp, []))
        for comp in self.enrichment.index
    ]

    df = pd.DataFrame(
        {
            "tau": self.tau,
            "eta_squared": self.eta_squared,
            "kruskal_epsilon_squared": self.kruskal_epsilon_squared,
            "top_cell_type": top_types,
            "significant_cell_types": sig_str,
        },
        index=self.enrichment.index,
    )
    return df

ComponentAlignmentResult dataclass

Alignment results for a single RISE component.

Parameters:

Name Type Description Default
component int | str

Component index or label.

required
enrichment Series

AUROC per cell type distinguishing that type from all others.

required
p_values Series

Empirical permutation p-values for enrichment (AUROC > null).

required
q_values Series

Benjamini-Hochberg FDR-adjusted p-values.

required
tau float

Tissue specificity index tau (uniqueness score in [0, 1]).

required
eta_squared float

Proportion of loading variance explained by cell type (in [0, 1]).

required
kruskal_epsilon_squared float

Non-parametric effect size epsilon-squared from Kruskal-Wallis.

required
significant_cell_types list[str]

List of cell types with significant enrichment (q <= alpha and AUROC > 0.5).

list()
alpha float

Significance threshold used for q-values.

0.05
Source code in scrise/annotation_alignment.py
@dataclass
class ComponentAlignmentResult:
    """Alignment results for a single RISE component.

    Parameters
    ----------
    component : int | str
        Component index or label.
    enrichment : pd.Series
        AUROC per cell type distinguishing that type from all others.
    p_values : pd.Series
        Empirical permutation p-values for enrichment (AUROC > null).
    q_values : pd.Series
        Benjamini-Hochberg FDR-adjusted p-values.
    tau : float
        Tissue specificity index tau (uniqueness score in [0, 1]).
    eta_squared : float
        Proportion of loading variance explained by cell type (in [0, 1]).
    kruskal_epsilon_squared : float
        Non-parametric effect size epsilon-squared from Kruskal-Wallis.
    significant_cell_types : list[str]
        List of cell types with significant enrichment (q <= alpha and AUROC > 0.5).
    alpha : float
        Significance threshold used for q-values.
    """

    component: int | str
    enrichment: pd.Series
    p_values: pd.Series
    q_values: pd.Series
    tau: float
    eta_squared: float
    kruskal_epsilon_squared: float
    significant_cell_types: list[str] = field(default_factory=list)
    alpha: float = 0.05

    def to_dict(self) -> dict[str, Any]:
        """Convert result to a dictionary."""
        return {
            "component": self.component,
            "enrichment": self.enrichment.to_dict(),
            "p_values": self.p_values.to_dict(),
            "q_values": self.q_values.to_dict(),
            "tau": self.tau,
            "eta_squared": self.eta_squared,
            "kruskal_epsilon_squared": self.kruskal_epsilon_squared,
            "significant_cell_types": list(self.significant_cell_types),
            "alpha": self.alpha,
        }

to_dict

to_dict() -> dict[str, Any]

Convert result to a dictionary.

Source code in scrise/annotation_alignment.py
def to_dict(self) -> dict[str, Any]:
    """Convert result to a dictionary."""
    return {
        "component": self.component,
        "enrichment": self.enrichment.to_dict(),
        "p_values": self.p_values.to_dict(),
        "q_values": self.q_values.to_dict(),
        "tau": self.tau,
        "eta_squared": self.eta_squared,
        "kruskal_epsilon_squared": self.kruskal_epsilon_squared,
        "significant_cell_types": list(self.significant_cell_types),
        "alpha": self.alpha,
    }

cell_type_alignment

cell_type_alignment(
    loadings: ndarray | Series,
    cell_types: Series | ndarray,
    signed: bool = False,
    n_permutations: int = 1000,
    alpha: float = 0.05,
    random_state: int | Generator | None = None,
    component_label: int | str = 1,
) -> ComponentAlignmentResult

Score alignment of a single component's cell loadings with annotated cell types.

Parameters:

Name Type Description Default
loadings ndarray | Series

Cell loading vector for one component (shape: n_cells,).

required
cell_types Series | ndarray

Cell type annotations (length: n_cells).

required
signed bool, optional (default: False)

If True, uses the absolute value of loadings (|loading|), which is appropriate for signed eigen-state projections (e.g. P_i @ B[:, r]).

False
n_permutations int, optional (default: 1000)

Number of label permutations to compute empirical p-values. If 0, p-values are computed using the asymptotic one-sided Mann-Whitney U test.

1000
alpha float, optional (default: 0.05)

FDR significance threshold for identifying enriched cell types.

0.05
random_state int | np.random.Generator | None, optional (default: None)

Random seed or Generator for permutation reproducibility.

None
component_label int | str, optional (default: 1)

Identifier for the component.

1

Returns:

Type Description
ComponentAlignmentResult

Dataclass containing AUROC enrichment, p-values, q-values, tau, eta^2, and significant cell types.

Source code in scrise/annotation_alignment.py
def cell_type_alignment(
    loadings: np.ndarray | pd.Series,
    cell_types: pd.Series | np.ndarray,
    signed: bool = False,
    n_permutations: int = 1000,
    alpha: float = 0.05,
    random_state: int | np.random.Generator | None = None,
    component_label: int | str = 1,
) -> ComponentAlignmentResult:
    """Score alignment of a single component's cell loadings with annotated cell types.

    Parameters
    ----------
    loadings : np.ndarray | pd.Series
        Cell loading vector for one component (shape: n_cells,).
    cell_types : pd.Series | np.ndarray
        Cell type annotations (length: n_cells).
    signed : bool, optional (default: False)
        If True, uses the absolute value of loadings (|loading|), which is appropriate
        for signed eigen-state projections (e.g. P_i @ B[:, r]).
    n_permutations : int, optional (default: 1000)
        Number of label permutations to compute empirical p-values. If 0, p-values
        are computed using the asymptotic one-sided Mann-Whitney U test.
    alpha : float, optional (default: 0.05)
        FDR significance threshold for identifying enriched cell types.
    random_state : int | np.random.Generator | None, optional (default: None)
        Random seed or Generator for permutation reproducibility.
    component_label : int | str, optional (default: 1)
        Identifier for the component.

    Returns
    -------
    ComponentAlignmentResult
        Dataclass containing AUROC enrichment, p-values, q-values, tau, eta^2,
        and significant cell types.
    """
    y = np.asarray(loadings, dtype=float)
    if signed:
        y = np.abs(y)

    codes, categories = _validate_and_encode_cell_types(cell_types)
    n_types = len(categories)
    n_cells = y.size

    if codes.size != n_cells:
        raise ValueError(
            f"Length mismatch: loadings has {n_cells} cells but cell_types has {codes.size}."
        )

    # Compute observed AUROC per cell type
    aurocs = compute_auroc_per_cell_type(y, codes, n_types)

    p_values = _component_p_values(
        y,
        codes,
        n_types,
        n_cells,
        aurocs,
        n_permutations,
        _as_generator(random_state),
    )

    # BH FDR correction across cell types for this component
    if n_types > 1:
        q_values = sp.false_discovery_control(p_values, method="bh")
    else:
        q_values = p_values.copy()

    # Scores
    tau = compute_tau(aurocs, baseline=0.0)
    eta2 = compute_eta_squared(y, codes, n_types)
    eps2 = compute_kruskal_epsilon_squared(y, codes, n_types)

    enrichment_series = pd.Series(aurocs, index=categories, name="AUROC")
    p_series = pd.Series(p_values, index=categories, name="p_value")
    q_series = pd.Series(q_values, index=categories, name="q_value")

    significant = [
        categories[k]
        for k in range(n_types)
        if q_values[k] <= alpha and aurocs[k] > 0.5
    ]

    return ComponentAlignmentResult(
        component=component_label,
        enrichment=enrichment_series,
        p_values=p_series,
        q_values=q_series,
        tau=tau,
        eta_squared=eta2,
        kruskal_epsilon_squared=eps2,
        significant_cell_types=significant,
        alpha=alpha,
    )

score_cell_type_alignment

score_cell_type_alignment(
    data: AnnData | ndarray | DataFrame,
    cell_types: Series | ndarray | str | None = None,
    signed: bool = False,
    projection_key: str = "weighted_projections",
    n_permutations: int = 1000,
    alpha: float = 0.05,
    random_state: int | Generator | None = None,
) -> CellTypeAlignmentResults

Score cell-type alignment across all RISE components.

Jointly calculates per-cell-type AUROC enrichment, empirical significance with Benjamini-Hochberg FDR correction across all (component x cell type) tests, uniqueness (tau), and combination alignment (eta^2).

Parameters:

Name Type Description Default
data AnnData | ndarray | DataFrame

AnnData containing fitted RISE results, or a matrix of cell loadings (shape: n_cells, n_components).

required
cell_types Series | ndarray | str | None

Cell type annotations. If data is AnnData and cell_types is a string (or None), looks up data.obs[cell_types] (defaults to 'cell_type' or 'CellType').

None
signed bool, optional (default: False)

If True, takes the absolute value of loadings (|loading|).

False
projection_key str, optional (default: "weighted_projections")

Key in data.obsm to extract loadings from when data is an AnnData. Defaults to 'weighted_projections', or falls back to 'projections'.

'weighted_projections'
n_permutations int, optional (default: 1000)

Number of permutations for null AUROC distribution.

1000
alpha float, optional (default: 0.05)

Significance threshold for FDR q-values.

0.05
random_state int | np.random.Generator | None, optional (default: None)

Random seed or Generator for permutations.

None

Returns:

Type Description
CellTypeAlignmentResults

Container with full results across all components.

Source code in scrise/annotation_alignment.py
def score_cell_type_alignment(
    data: anndata.AnnData | np.ndarray | pd.DataFrame,
    cell_types: pd.Series | np.ndarray | str | None = None,
    signed: bool = False,
    projection_key: str = "weighted_projections",
    n_permutations: int = 1000,
    alpha: float = 0.05,
    random_state: int | np.random.Generator | None = None,
) -> CellTypeAlignmentResults:
    """Score cell-type alignment across all RISE components.

    Jointly calculates per-cell-type AUROC enrichment, empirical significance with
    Benjamini-Hochberg FDR correction across all (component x cell type) tests,
    uniqueness (tau), and combination alignment (eta^2).

    Parameters
    ----------
    data : anndata.AnnData | np.ndarray | pd.DataFrame
        AnnData containing fitted RISE results, or a matrix of cell loadings
        (shape: n_cells, n_components).
    cell_types : pd.Series | np.ndarray | str | None, optional
        Cell type annotations. If data is AnnData and cell_types is a string (or None),
        looks up data.obs[cell_types] (defaults to 'cell_type' or 'CellType').
    signed : bool, optional (default: False)
        If True, takes the absolute value of loadings (|loading|).
    projection_key : str, optional (default: "weighted_projections")
        Key in data.obsm to extract loadings from when data is an AnnData.
        Defaults to 'weighted_projections', or falls back to 'projections'.
    n_permutations : int, optional (default: 1000)
        Number of permutations for null AUROC distribution.
    alpha : float, optional (default: 0.05)
        Significance threshold for FDR q-values.
    random_state : int | np.random.Generator | None, optional (default: None)
        Random seed or Generator for permutations.

    Returns
    -------
    CellTypeAlignmentResults
        Container with full results across all components.
    """
    loadings_matrix, cell_type_series = _resolve_loadings_and_cell_types(
        data, cell_types, projection_key
    )

    if loadings_matrix.ndim == 1:
        loadings_matrix = loadings_matrix[:, np.newaxis]

    n_cells, n_comps = loadings_matrix.shape
    codes, categories = _validate_and_encode_cell_types(cell_type_series)
    n_types = len(categories)

    if codes.size != n_cells:
        raise ValueError(
            f"Length mismatch: data has {n_cells} cells but cell_types has {codes.size}."
        )

    rng = _as_generator(random_state)

    component_labels = [i + 1 for i in range(n_comps)]
    aurocs_mat, p_vals_mat, tau_vec, eta2_vec, eps2_vec = _component_metrics(
        loadings_matrix, codes, n_types, signed, n_permutations, rng
    )

    # Joint BH FDR correction across all (component x cell_type) tests
    if aurocs_mat.size > 1:
        q_vals_mat = sp.false_discovery_control(
            p_vals_mat.ravel(), method="bh"
        ).reshape(p_vals_mat.shape)
    else:
        q_vals_mat = p_vals_mat.copy()

    enrichment_df = pd.DataFrame(aurocs_mat, index=component_labels, columns=categories)
    p_values_df = pd.DataFrame(p_vals_mat, index=component_labels, columns=categories)
    q_values_df = pd.DataFrame(q_vals_mat, index=component_labels, columns=categories)
    tau_series = pd.Series(tau_vec, index=component_labels, name="tau")
    eta2_series = pd.Series(eta2_vec, index=component_labels, name="eta_squared")
    eps2_series = pd.Series(
        eps2_vec, index=component_labels, name="kruskal_epsilon_squared"
    )

    results_list: list[ComponentAlignmentResult] = []
    sig_dict: dict[int | str, list[str]] = {}

    for comp_idx, comp_lbl in enumerate(component_labels):
        sig_types = [
            categories[k]
            for k in range(n_types)
            if q_vals_mat[comp_idx, k] <= alpha and aurocs_mat[comp_idx, k] > 0.5
        ]
        sig_dict[comp_lbl] = sig_types
        res = ComponentAlignmentResult(
            component=comp_lbl,
            enrichment=enrichment_df.loc[comp_lbl],
            p_values=p_values_df.loc[comp_lbl],
            q_values=q_values_df.loc[comp_lbl],
            tau=tau_vec[comp_idx],
            eta_squared=eta2_vec[comp_idx],
            kruskal_epsilon_squared=eps2_vec[comp_idx],
            significant_cell_types=sig_types,
            alpha=alpha,
        )
        results_list.append(res)

    return CellTypeAlignmentResults(
        results=results_list,
        enrichment=enrichment_df,
        p_values=p_values_df,
        q_values=q_values_df,
        tau=tau_series,
        eta_squared=eta2_series,
        kruskal_epsilon_squared=eps2_series,
        significant_cell_types=sig_dict,
        alpha=alpha,
    )

Alignment Statistics

scrise.alignment_stats

compute_auroc_per_cell_type

compute_auroc_per_cell_type(
    loadings: ndarray,
    cell_type_codes: ndarray,
    n_types: int,
) -> np.ndarray

Compute AUROC for each cell type vs all other cells.

Parameters:

Name Type Description Default
loadings ndarray

1D array of cell loadings of shape (n_cells,).

required
cell_type_codes ndarray

1D array of integer cell type assignments in [0, n_types - 1].

required
n_types int

Total number of unique cell types.

required

Returns:

Type Description
ndarray

1D array of AUROC values for each cell type of shape (n_types,).

Source code in scrise/alignment_stats.py
def compute_auroc_per_cell_type(
    loadings: np.ndarray,
    cell_type_codes: np.ndarray,
    n_types: int,
) -> np.ndarray:
    """Compute AUROC for each cell type vs all other cells.

    Parameters
    ----------
    loadings : np.ndarray
        1D array of cell loadings of shape (n_cells,).
    cell_type_codes : np.ndarray
        1D array of integer cell type assignments in [0, n_types - 1].
    n_types : int
        Total number of unique cell types.

    Returns
    -------
    np.ndarray
        1D array of AUROC values for each cell type of shape (n_types,).
    """
    n_cells = loadings.size
    if n_cells == 0 or n_types <= 1:
        return np.full(n_types, 0.5, dtype=float)

    ranks = sp.rankdata(loadings, method="average")
    counts = np.bincount(cell_type_codes, minlength=n_types).astype(float)
    rank_sums = np.bincount(cell_type_codes, weights=ranks, minlength=n_types).astype(
        float
    )

    aurocs = np.full(n_types, 0.5, dtype=float)
    for k in range(n_types):
        n1 = counts[k]
        n0 = n_cells - n1
        if n1 > 0 and n0 > 0:
            u1 = rank_sums[k] - n1 * (n1 + 1.0) / 2.0
            aurocs[k] = u1 / (n1 * n0)

    return aurocs

compute_tau

compute_tau(
    enrichment_scores: ndarray | Series,
    baseline: float = 0.0,
) -> float

Compute tissue specificity index tau (Yanai et al., 2005).

tau = sum(1 - x_hat) / (n_types - 1), where x_hat = x / max(x).

When computed over AUROC enrichment values: - tau -> 1 indicates the component is specific/private to a single cell type. - tau -> 0 indicates the component is evenly distributed across cell types.

Parameters:

Name Type Description Default
enrichment_scores ndarray | Series

1D array of enrichment values (e.g. AUROC) across cell types.

required
baseline float, optional (default: 0.0)

Baseline value subtracted before calculating tau. Values below baseline are clipped to 0.

0.0

Returns:

Type Description
float

Tau index in [0.0, 1.0].

Source code in scrise/alignment_stats.py
def compute_tau(
    enrichment_scores: np.ndarray | pd.Series,
    baseline: float = 0.0,
) -> float:
    """Compute tissue specificity index tau (Yanai et al., 2005).

    tau = sum(1 - x_hat) / (n_types - 1), where x_hat = x / max(x).

    When computed over AUROC enrichment values:
    - tau -> 1 indicates the component is specific/private to a single cell type.
    - tau -> 0 indicates the component is evenly distributed across cell types.

    Parameters
    ----------
    enrichment_scores : np.ndarray | pd.Series
        1D array of enrichment values (e.g. AUROC) across cell types.
    baseline : float, optional (default: 0.0)
        Baseline value subtracted before calculating tau. Values below baseline are clipped to 0.

    Returns
    -------
    float
        Tau index in [0.0, 1.0].
    """
    x = np.asarray(enrichment_scores, dtype=float)
    n = x.size
    if n <= 1:
        return 0.0

    if baseline > 0.0:
        x = np.maximum(0.0, x - baseline)

    max_val = np.nanmax(x)
    if max_val <= 0.0 or not np.isfinite(max_val):
        return 0.0

    x_hat = x / max_val
    tau = np.sum(1.0 - x_hat) / (n - 1.0)
    return float(np.clip(tau, 0.0, 1.0))

compute_eta_squared

compute_eta_squared(
    loadings: ndarray,
    cell_type_codes: ndarray,
    n_types: int,
) -> float

Compute omnibus eta-squared (loading ~ cell_type).

Proportion of loading variance explained by cell-type identity.

Parameters:

Name Type Description Default
loadings ndarray

1D array of cell loadings (shape: n_cells,).

required
cell_type_codes ndarray

1D array of integer cell type assignments (shape: n_cells,).

required
n_types int

Number of unique cell types.

required

Returns:

Type Description
float

Eta-squared value in [0.0, 1.0].

Source code in scrise/alignment_stats.py
def compute_eta_squared(
    loadings: np.ndarray,
    cell_type_codes: np.ndarray,
    n_types: int,
) -> float:
    """Compute omnibus eta-squared (loading ~ cell_type).

    Proportion of loading variance explained by cell-type identity.

    Parameters
    ----------
    loadings : np.ndarray
        1D array of cell loadings (shape: n_cells,).
    cell_type_codes : np.ndarray
        1D array of integer cell type assignments (shape: n_cells,).
    n_types : int
        Number of unique cell types.

    Returns
    -------
    float
        Eta-squared value in [0.0, 1.0].
    """
    y = np.asarray(loadings, dtype=float)
    n_cells = y.size
    if n_cells <= 1 or n_types <= 1:
        return 0.0

    y_mean = np.mean(y)
    ss_total = np.sum((y - y_mean) ** 2)
    if ss_total <= 1e-12:
        return 0.0

    counts = np.bincount(cell_type_codes, minlength=n_types).astype(float)
    valid = counts > 0
    sums = np.bincount(cell_type_codes, weights=y, minlength=n_types)

    means = np.zeros(n_types, dtype=float)
    means[valid] = sums[valid] / counts[valid]

    ss_between = np.sum(counts[valid] * (means[valid] - y_mean) ** 2)
    eta2 = ss_between / ss_total
    return float(np.clip(eta2, 0.0, 1.0))

compute_kruskal_epsilon_squared

compute_kruskal_epsilon_squared(
    loadings: ndarray,
    cell_type_codes: ndarray,
    n_types: int,
) -> float

Compute Kruskal-Wallis epsilon-squared effect size.

Non-parametric measure of association between loading and cell type.

Parameters:

Name Type Description Default
loadings ndarray

1D array of cell loadings (shape: n_cells,).

required
cell_type_codes ndarray

1D array of integer cell type assignments (shape: n_cells,).

required
n_types int

Number of unique cell types.

required

Returns:

Type Description
float

Epsilon-squared value in [0.0, 1.0].

Source code in scrise/alignment_stats.py
def compute_kruskal_epsilon_squared(
    loadings: np.ndarray,
    cell_type_codes: np.ndarray,
    n_types: int,
) -> float:
    """Compute Kruskal-Wallis epsilon-squared effect size.

    Non-parametric measure of association between loading and cell type.

    Parameters
    ----------
    loadings : np.ndarray
        1D array of cell loadings (shape: n_cells,).
    cell_type_codes : np.ndarray
        1D array of integer cell type assignments (shape: n_cells,).
    n_types : int
        Number of unique cell types.

    Returns
    -------
    float
        Epsilon-squared value in [0.0, 1.0].
    """
    y = np.asarray(loadings, dtype=float)
    n_cells = y.size
    if n_cells <= 1 or n_types <= 1:
        return 0.0

    groups = [
        y[cell_type_codes == k]
        for k in range(n_types)
        if np.sum(cell_type_codes == k) > 0
    ]
    if len(groups) <= 1:
        return 0.0

    try:
        stat, _ = sp.kruskal(*groups)
        if not np.isfinite(stat) or stat < 0:
            return 0.0
        eps2 = stat / (n_cells - 1.0)
        return float(np.clip(eps2, 0.0, 1.0))
    except (ValueError, ZeroDivisionError):
        return 0.0

Quantization & Compression

scrise.opq

Optimized Product Quantization (OPQ) for projecting and compressing factor matrices.

OPQQuantizer

Optimized Product Quantizer for compressing projection matrices.

Parameters:

Name Type Description Default
M int

Number of sub-quantizers (sub-vector partitions).

required
n_bits int, optional (default: 8)

Number of bits per sub-quantizer. 8 bits corresponds to 256 centroids per sub-space.

8
n_iter int, optional (default: 5)

Number of alternating OPQ optimization iterations for the rotation matrix.

5
random_state int, optional (default: 42)

Random seed for reproducibility.

42
Source code in scrise/opq.py
class OPQQuantizer:
    """Optimized Product Quantizer for compressing projection matrices.

    Parameters
    ----------
    M : int
        Number of sub-quantizers (sub-vector partitions).
    n_bits : int, optional (default: 8)
        Number of bits per sub-quantizer. 8 bits corresponds to 256 centroids per sub-space.
    n_iter : int, optional (default: 5)
        Number of alternating OPQ optimization iterations for the rotation matrix.
    random_state : int, optional (default: 42)
        Random seed for reproducibility.
    """

    def __init__(
        self,
        M: int,
        n_bits: int = 8,
        n_iter: int = 5,
        random_state: int = 42,
    ):
        self.M = M
        self.n_bits = n_bits
        self.n_iter = n_iter
        self.random_state = random_state
        self.R: np.ndarray | None = None
        self.centroids_cat: np.ndarray | None = None
        self.sub_dims: np.ndarray | None = None

    def fit(self, X: np.ndarray, sample_size: int = 50000) -> "OPQQuantizer":
        """Fit rotation matrix and sub-quantizer codebooks on data matrix X."""
        X_arr = np.asarray(X, dtype=np.float32)
        N, D = X_arr.shape
        M = min(self.M, D)
        self.M = M

        # Partition dimensions across M sub-vectors
        sub_dims = [D // M + (1 if i < D % M else 0) for i in range(M)]
        self.sub_dims = np.array(sub_dims, dtype=np.int32)
        dim_offsets = np.cumsum([0] + sub_dims)
        K_centroids = min(2**self.n_bits, N)

        # Subsample for training if N > sample_size
        if N > sample_size:
            rng = np.random.default_rng(self.random_state)
            idx = rng.choice(N, size=sample_size, replace=False)
            train_X = X_arr[idx]
        else:
            train_X = X_arr

        # Initialize rotation matrix
        R = np.eye(D, dtype=np.float32)

        # Alternating OPQ iterations
        for _ in range(self.n_iter):
            X_rot = train_X @ R
            X_hat_rot = np.zeros_like(X_rot)
            for m in range(M):
                ds, de = dim_offsets[m], dim_offsets[m + 1]
                sub = X_rot[:, ds:de]
                if sub.shape[1] == 1:
                    km = KMeans(
                        n_clusters=K_centroids,
                        random_state=self.random_state,
                        n_init=1,
                        max_iter=20,
                    ).fit(sub)
                else:
                    km = MiniBatchKMeans(
                        n_clusters=K_centroids,
                        random_state=self.random_state,
                        batch_size=min(2048, len(train_X)),
                        n_init=1,
                        max_iter=20,
                    ).fit(sub)
                X_hat_rot[:, ds:de] = km.cluster_centers_[km.labels_]

            # Update orthogonal rotation matrix via Procrustes SVD
            U, _, Vt = np.linalg.svd(train_X.T @ X_hat_rot)
            R = (U @ Vt).astype(np.float32)

        # Final codebook fitting on rotated training data
        X_rot_train = train_X @ R
        centroids_list = []
        for m in range(M):
            ds, de = dim_offsets[m], dim_offsets[m + 1]
            sub = X_rot_train[:, ds:de]
            if sub.shape[1] == 1:
                km = KMeans(
                    n_clusters=K_centroids,
                    random_state=self.random_state,
                    n_init=2,
                    max_iter=50,
                ).fit(sub)
            else:
                km = MiniBatchKMeans(
                    n_clusters=K_centroids,
                    random_state=self.random_state,
                    batch_size=min(4096, len(train_X)),
                    n_init=3,
                    max_iter=50,
                ).fit(sub)
            centroids_list.append(km.cluster_centers_.astype(np.float32))

        self.R = R
        self.centroids_cat = np.concatenate(centroids_list, axis=1).astype(np.float32)
        return self

    def encode(self, X: np.ndarray) -> np.ndarray:
        """Encode data matrix X into uint8 sub-quantizer centroid indices."""
        if self.R is None or self.centroids_cat is None or self.sub_dims is None:
            raise ValueError("OPQQuantizer has not been fitted yet.")

        X_arr = np.asarray(X, dtype=np.float32)
        N, _ = X_arr.shape
        M = self.M
        dim_offsets = np.cumsum([0] + list(self.sub_dims))
        X_rot = X_arr @ self.R
        codes = np.zeros((N, M), dtype=np.uint8)

        for m in range(M):
            ds, de = dim_offsets[m], dim_offsets[m + 1]
            sub = X_rot[:, ds:de]
            c = self.centroids_cat[:, ds:de]
            # Vectorized nearest centroid computation: dist^2 = ||sub||^2 - 2 sub @ c.T + ||c||^2
            sub_sq = np.sum(sub**2, axis=1, keepdims=True)
            c_sq = np.sum(c**2, axis=1, keepdims=True).T
            dists = sub_sq - 2.0 * (sub @ c.T) + c_sq
            codes[:, m] = np.argmin(dists, axis=1).astype(np.uint8)

        return codes

    def decode(self, codes: np.ndarray) -> np.ndarray:
        """Decode uint8 sub-quantizer codes back to the continuous feature space."""
        if self.R is None or self.centroids_cat is None or self.sub_dims is None:
            raise ValueError("OPQQuantizer has not been fitted yet.")

        codes_arr = np.asarray(codes, dtype=np.uint8)
        N, M = codes_arr.shape
        D = self.R.shape[0]
        dim_offsets = np.cumsum([0] + list(self.sub_dims))
        X_recon_rot = np.zeros((N, D), dtype=np.float32)

        for m in range(M):
            ds, de = dim_offsets[m], dim_offsets[m + 1]
            c = self.centroids_cat[:, ds:de]
            X_recon_rot[:, ds:de] = c[codes_arr[:, m]]

        return (X_recon_rot @ self.R.T).astype(np.float32)

    def fit_transform(
        self, X: np.ndarray, sample_size: int = 50000
    ) -> tuple[np.ndarray, np.ndarray, float]:
        """Fit OPQ, encode X to codes, reconstruct, and compute R^2."""
        self.fit(X, sample_size=sample_size)
        codes = self.encode(X)
        recon = self.decode(codes)
        r2 = float(r2_score(X, recon))
        return codes, recon, r2

    @classmethod
    def from_saved(
        cls,
        R: np.ndarray,
        centroids_cat: np.ndarray,
        sub_dims: np.ndarray,
        n_bits: int = 8,
    ) -> "OPQQuantizer":
        """Instantiate a fitted OPQQuantizer from saved parameters."""
        quantizer = cls(M=len(sub_dims), n_bits=n_bits)
        quantizer.R = np.asarray(R, dtype=np.float32)
        quantizer.centroids_cat = np.asarray(centroids_cat, dtype=np.float32)
        quantizer.sub_dims = np.asarray(sub_dims, dtype=np.int32)
        return quantizer

decode

decode(codes: ndarray) -> np.ndarray

Decode uint8 sub-quantizer codes back to the continuous feature space.

Source code in scrise/opq.py
def decode(self, codes: np.ndarray) -> np.ndarray:
    """Decode uint8 sub-quantizer codes back to the continuous feature space."""
    if self.R is None or self.centroids_cat is None or self.sub_dims is None:
        raise ValueError("OPQQuantizer has not been fitted yet.")

    codes_arr = np.asarray(codes, dtype=np.uint8)
    N, M = codes_arr.shape
    D = self.R.shape[0]
    dim_offsets = np.cumsum([0] + list(self.sub_dims))
    X_recon_rot = np.zeros((N, D), dtype=np.float32)

    for m in range(M):
        ds, de = dim_offsets[m], dim_offsets[m + 1]
        c = self.centroids_cat[:, ds:de]
        X_recon_rot[:, ds:de] = c[codes_arr[:, m]]

    return (X_recon_rot @ self.R.T).astype(np.float32)

encode

encode(X: ndarray) -> np.ndarray

Encode data matrix X into uint8 sub-quantizer centroid indices.

Source code in scrise/opq.py
def encode(self, X: np.ndarray) -> np.ndarray:
    """Encode data matrix X into uint8 sub-quantizer centroid indices."""
    if self.R is None or self.centroids_cat is None or self.sub_dims is None:
        raise ValueError("OPQQuantizer has not been fitted yet.")

    X_arr = np.asarray(X, dtype=np.float32)
    N, _ = X_arr.shape
    M = self.M
    dim_offsets = np.cumsum([0] + list(self.sub_dims))
    X_rot = X_arr @ self.R
    codes = np.zeros((N, M), dtype=np.uint8)

    for m in range(M):
        ds, de = dim_offsets[m], dim_offsets[m + 1]
        sub = X_rot[:, ds:de]
        c = self.centroids_cat[:, ds:de]
        # Vectorized nearest centroid computation: dist^2 = ||sub||^2 - 2 sub @ c.T + ||c||^2
        sub_sq = np.sum(sub**2, axis=1, keepdims=True)
        c_sq = np.sum(c**2, axis=1, keepdims=True).T
        dists = sub_sq - 2.0 * (sub @ c.T) + c_sq
        codes[:, m] = np.argmin(dists, axis=1).astype(np.uint8)

    return codes

fit

fit(X: ndarray, sample_size: int = 50000) -> OPQQuantizer

Fit rotation matrix and sub-quantizer codebooks on data matrix X.

Source code in scrise/opq.py
def fit(self, X: np.ndarray, sample_size: int = 50000) -> "OPQQuantizer":
    """Fit rotation matrix and sub-quantizer codebooks on data matrix X."""
    X_arr = np.asarray(X, dtype=np.float32)
    N, D = X_arr.shape
    M = min(self.M, D)
    self.M = M

    # Partition dimensions across M sub-vectors
    sub_dims = [D // M + (1 if i < D % M else 0) for i in range(M)]
    self.sub_dims = np.array(sub_dims, dtype=np.int32)
    dim_offsets = np.cumsum([0] + sub_dims)
    K_centroids = min(2**self.n_bits, N)

    # Subsample for training if N > sample_size
    if N > sample_size:
        rng = np.random.default_rng(self.random_state)
        idx = rng.choice(N, size=sample_size, replace=False)
        train_X = X_arr[idx]
    else:
        train_X = X_arr

    # Initialize rotation matrix
    R = np.eye(D, dtype=np.float32)

    # Alternating OPQ iterations
    for _ in range(self.n_iter):
        X_rot = train_X @ R
        X_hat_rot = np.zeros_like(X_rot)
        for m in range(M):
            ds, de = dim_offsets[m], dim_offsets[m + 1]
            sub = X_rot[:, ds:de]
            if sub.shape[1] == 1:
                km = KMeans(
                    n_clusters=K_centroids,
                    random_state=self.random_state,
                    n_init=1,
                    max_iter=20,
                ).fit(sub)
            else:
                km = MiniBatchKMeans(
                    n_clusters=K_centroids,
                    random_state=self.random_state,
                    batch_size=min(2048, len(train_X)),
                    n_init=1,
                    max_iter=20,
                ).fit(sub)
            X_hat_rot[:, ds:de] = km.cluster_centers_[km.labels_]

        # Update orthogonal rotation matrix via Procrustes SVD
        U, _, Vt = np.linalg.svd(train_X.T @ X_hat_rot)
        R = (U @ Vt).astype(np.float32)

    # Final codebook fitting on rotated training data
    X_rot_train = train_X @ R
    centroids_list = []
    for m in range(M):
        ds, de = dim_offsets[m], dim_offsets[m + 1]
        sub = X_rot_train[:, ds:de]
        if sub.shape[1] == 1:
            km = KMeans(
                n_clusters=K_centroids,
                random_state=self.random_state,
                n_init=2,
                max_iter=50,
            ).fit(sub)
        else:
            km = MiniBatchKMeans(
                n_clusters=K_centroids,
                random_state=self.random_state,
                batch_size=min(4096, len(train_X)),
                n_init=3,
                max_iter=50,
            ).fit(sub)
        centroids_list.append(km.cluster_centers_.astype(np.float32))

    self.R = R
    self.centroids_cat = np.concatenate(centroids_list, axis=1).astype(np.float32)
    return self

fit_transform

fit_transform(
    X: ndarray, sample_size: int = 50000
) -> tuple[np.ndarray, np.ndarray, float]

Fit OPQ, encode X to codes, reconstruct, and compute R^2.

Source code in scrise/opq.py
def fit_transform(
    self, X: np.ndarray, sample_size: int = 50000
) -> tuple[np.ndarray, np.ndarray, float]:
    """Fit OPQ, encode X to codes, reconstruct, and compute R^2."""
    self.fit(X, sample_size=sample_size)
    codes = self.encode(X)
    recon = self.decode(codes)
    r2 = float(r2_score(X, recon))
    return codes, recon, r2

from_saved classmethod

from_saved(
    R: ndarray,
    centroids_cat: ndarray,
    sub_dims: ndarray,
    n_bits: int = 8,
) -> OPQQuantizer

Instantiate a fitted OPQQuantizer from saved parameters.

Source code in scrise/opq.py
@classmethod
def from_saved(
    cls,
    R: np.ndarray,
    centroids_cat: np.ndarray,
    sub_dims: np.ndarray,
    n_bits: int = 8,
) -> "OPQQuantizer":
    """Instantiate a fitted OPQQuantizer from saved parameters."""
    quantizer = cls(M=len(sub_dims), n_bits=n_bits)
    quantizer.R = np.asarray(R, dtype=np.float32)
    quantizer.centroids_cat = np.asarray(centroids_cat, dtype=np.float32)
    quantizer.sub_dims = np.asarray(sub_dims, dtype=np.int32)
    return quantizer

find_optimal_opq

find_optimal_opq(
    P: ndarray,
    fidelity_threshold: float = 0.99,
    random_state: int = 42,
) -> tuple[OPQQuantizer, np.ndarray, float]

Find the smallest number of sub-quantizers M achieving R^2 >= fidelity_threshold.

Parameters:

Name Type Description Default
P ndarray

Projection matrix of shape (N, D).

required
fidelity_threshold float, optional (default: 0.99)

Target R^2 reconstruction accuracy.

0.99
random_state int, optional (default: 42)

Random seed.

42

Returns:

Type Description
tuple of (OPQQuantizer, np.ndarray, float)

(quantizer, codes, r2)

Source code in scrise/opq.py
def find_optimal_opq(
    P: np.ndarray,
    fidelity_threshold: float = 0.99,
    random_state: int = 42,
) -> tuple[OPQQuantizer, np.ndarray, float]:
    """Find the smallest number of sub-quantizers M achieving R^2 >= fidelity_threshold.

    Parameters
    ----------
    P : np.ndarray
        Projection matrix of shape (N, D).
    fidelity_threshold : float, optional (default: 0.99)
        Target R^2 reconstruction accuracy.
    random_state : int, optional (default: 42)
        Random seed.

    Returns
    -------
    tuple of (OPQQuantizer, np.ndarray, float)
        (quantizer, codes, r2)
    """
    P_arr = np.asarray(P, dtype=np.float32)
    _, D = P_arr.shape

    # Candidate M values to evaluate (increasing order for compression)
    candidate_Ms = sorted(
        {m for m in (1, 2, 4, 5, 8, 10, 12, 15, 20, 25, 30, D) if m <= D}
    )

    best_quantizer = None
    best_codes = None
    best_r2 = -1.0

    for M in candidate_Ms:
        quantizer = OPQQuantizer(M=M, random_state=random_state)
        codes, _, r2 = quantizer.fit_transform(P_arr)

        if r2 > best_r2:
            best_quantizer = quantizer
            best_codes = codes
            best_r2 = r2

        if r2 >= fidelity_threshold:
            return quantizer, codes, r2

    # Fallback to highest fidelity achieved: the loop always runs at least
    # once (candidate_Ms is never empty), so both are guaranteed to be set.
    assert best_quantizer is not None
    assert best_codes is not None
    return best_quantizer, best_codes, best_r2

Preprocessing

parafac2.normalize

Dataset preprocessing and normalization utilities for PARAFAC2 analysis.

This module provides functions to filter, normalize, and annotate single-cell gene expression datasets stored in AnnData objects prior to PARAFAC2 matrix factorization.

prepare_dataset

prepare_dataset(
    X: AnnData, condition_name: str, geneThreshold: float
) -> anndata.AnnData

Preprocess and normalize an AnnData dataset for PARAFAC2 factorization.

Performs quality control filtering of low-count cells and low-expression genes, normalizes total cell counts and gene sums, applies a log10 transformation, and computes metadata required by PARAFAC2 (condition indices and gene means).

Parameters:

Name Type Description Default
X AnnData

Input single-cell dataset with raw count matrix stored in X.X (must be a sparse matrix with non-negative values).

required
condition_name str

Column name in X.obs identifying the sample or experimental condition grouping for each cell.

required
geneThreshold float

Minimum threshold fraction for gene inclusion. Genes with total counts less than geneThreshold * total_cells are filtered out.

required

Returns:

Type Description
AnnData

A filtered and normalized copy of the AnnData object. Contains the log-transformed normalized counts in X.X, integer condition codes in X.obs["condition_unique_idxs"], and per-gene mean expression values in X.var["means"].

Source code in parafac2/normalize.py
def prepare_dataset(
    X: anndata.AnnData, condition_name: str, geneThreshold: float
) -> anndata.AnnData:
    """Preprocess and normalize an AnnData dataset for PARAFAC2 factorization.

    Performs quality control filtering of low-count cells and low-expression
    genes, normalizes total cell counts and gene sums, applies a log10
    transformation, and computes metadata required by PARAFAC2 (condition
    indices and gene means).

    Parameters
    ----------
    X : anndata.AnnData
        Input single-cell dataset with raw count matrix stored in ``X.X``
        (must be a sparse matrix with non-negative values).
    condition_name : str
        Column name in ``X.obs`` identifying the sample or experimental
        condition grouping for each cell.
    geneThreshold : float
        Minimum threshold fraction for gene inclusion. Genes with total counts
        less than ``geneThreshold * total_cells`` are filtered out.

    Returns
    -------
    anndata.AnnData
        A filtered and normalized copy of the AnnData object. Contains the
        log-transformed normalized counts in ``X.X``, integer condition
        codes in ``X.obs["condition_unique_idxs"]``, and per-gene mean
        expression values in ``X.var["means"]``.
    """
    assert issparse(X.X)
    X_X_raw = cast("csr_array", X.X)
    assert np.amin(X_X_raw.data) >= 0.0

    # Filter out genes with too few reads, and cells with fewer than 10 counts
    cell_mask = np.ravel(X_X_raw.sum(axis=1)) > 10
    gene_mask = np.ravel(X_X_raw.sum(axis=0)) > (geneThreshold * X_X_raw.shape[0])

    if not cell_mask.any():
        raise ValueError(
            "prepare_dataset: no cells passed the count filter (every cell's "
            "total count is <= 10); check the input data."
        )
    if not gene_mask.any():
        raise ValueError(
            "prepare_dataset: no genes passed the count filter (every gene's "
            f"total count is <= {geneThreshold} * n_cells); check the input "
            "data or lower geneThreshold."
        )

    # Subset and materialize actual AnnData object before modifying X.X
    if cell_mask.all() and gene_mask.all():
        X = X.copy()
    else:
        X = X[cell_mask, gene_mask].copy()

    # Convert subset to csr_array and float32 data
    X.X = csr_array(X.X)
    X_X = cast("csr_array", X.X)

    if X_X.dtype != np.float32:
        X_X.data = X_X.data.astype(np.float32)

    ## Normalize total counts per cell
    # Keep the counts on a reasonable scale to avoid accuracy issues
    counts_per_cell = np.ravel(X_X.sum(axis=1)).astype(np.float32, copy=False)
    counts_per_cell /= np.median(counts_per_cell)
    # In-place CSR row scaling
    X_X.data /= np.repeat(counts_per_cell, np.diff(X_X.indptr))

    # Scale genes by sum, in-place CSR column scaling
    gene_sums = np.ravel(X_X.sum(axis=0)).astype(np.float32, copy=False)
    X_X.data /= gene_sums[X_X.indices]

    # Transform values in-place to avoid nnz-sized temporaries
    X_X.data *= np.float32(1000.0)
    X_X.data += np.float32(1.0)
    np.log10(X_X.data, out=X_X.data)

    # Get the indices for subsetting the data
    X.obs["condition_unique_idxs"] = pd.Categorical(X.obs[condition_name]).codes

    # Pre-calculate gene means
    X.var["means"] = np.ravel(X_X.mean(axis=0))

    return X

Visualization Functions

General Plotting

scrise.plotting.general

plot_r2x

plot_r2x(data, rank_vec, ax: Axes, compress='auto')

Plot variance explained (R²X) for RISE and PCA across different ranks.

This visualization helps determine the optimal number of components by showing how variance explained increases with rank. The elbow point where the curve flattens indicates a good balance between model complexity and explanatory power.

Parameters:

Name Type Description Default
data AnnData

Preprocessed AnnData object containing single-cell RNA-seq data. Must have X.obs["condition_unique_idxs"] for RISE decomposition.

required
rank_vec array-like of int

Array of rank values to test (e.g., [1, 5, 10, 15, 20, 25, 30]). Each rank represents a different number of components.

required
ax Axes

Matplotlib axes object to plot on.

required
compress int | tuple[int, int | None] | str | bool | None

CANDELINC compression mode passed to rise_pca_r2x for each rank's fit. Defaults to "auto". Pass None/False to fall back to exact ALS.

'auto'
Source code in scrise/plotting/general.py
def plot_r2x(data, rank_vec, ax: Axes, compress="auto"):
    """Plot variance explained (R²X) for RISE and PCA across different ranks.

    This visualization helps determine the optimal number of components by showing
    how variance explained increases with rank. The elbow point where the curve
    flattens indicates a good balance between model complexity and explanatory power.

    Parameters
    ----------
    data : anndata.AnnData
        Preprocessed AnnData object containing single-cell RNA-seq data.
        Must have X.obs["condition_unique_idxs"] for RISE decomposition.
    rank_vec : array-like of int
        Array of rank values to test (e.g., [1, 5, 10, 15, 20, 25, 30]).
        Each rank represents a different number of components.
    ax : matplotlib.axes.Axes
        Matplotlib axes object to plot on.
    compress : int | tuple[int, int | None] | str | bool | None, optional
        CANDELINC compression mode passed to ``rise_pca_r2x`` for each rank's
        fit. Defaults to ``"auto"``. Pass None/False to fall back to exact
        ALS.
    """
    r2xError = rise_pca_r2x(data, rank_vec, compress=compress)
    labelNames = ["Fit: RISE", "Fit: PCA"]
    colorDecomp = ["r", "b"]
    markerShape = ["o", "o"]
    for i in range(2):
        ax.scatter(
            rank_vec,
            r2xError[i],
            label=labelNames[i],
            marker=markerShape[i],
            c=colorDecomp[i],
            s=30.0,
        )
    ax.set(
        ylabel="Variance Explained",
        xlabel="Number of Components",
        xticks=np.linspace(0, rank_vec[-1], num=6, dtype=int),
        yticks=np.linspace(
            0, np.max(np.append(r2xError[0], r2xError[1])) + 0.01, num=5
        ),
    )
    ax.legend()

Factor Plotting

scrise.plotting.factors

plot_condition_factors

plot_condition_factors(
    data: AnnData,
    ax: Axes,
    cond: str = "Condition",
    log_transform: bool = True,
    cond_group_labels: Series | None = None,
    ThomsonNorm: bool = False,
    color_key=None,
    group_cond: bool = False,
    control_pattern: str | None = None,
    control_conditions: Sequence[str] | None = None,
)

Plot condition factors as a heatmap showing how conditions contribute to components.

This visualization shows how each experimental condition (rows) contributes to each RISE component (columns). High values indicate strong association between a condition and a component's pattern. Log transformation and normalization help reveal relative differences across conditions.

Parameters:

Name Type Description Default
data AnnData

AnnData object with RISE decomposition results. Must contain: - data.uns["Pf2_A"]: Condition factors (n_conditions, rank) - data.obs[cond]: Condition labels for each cell

required
ax Axes

Matplotlib axes object to plot on.

required
cond str, optional (default: "Condition")

Name of column in data.obs containing condition labels.

'Condition'
log_transform bool, optional (default: True)

If True, applies log10 transformation to condition factors before plotting. This helps visualize differences when values span orders of magnitude.

True
cond_group_labels pandas.Series, optional (default: None)

Series mapping conditions to group labels for colored row annotations. Useful for grouping related conditions (e.g., drug classes, patient cohorts).

None
ThomsonNorm bool, optional (default: False)

If True, normalizes factors using only control conditions (those containing 'CTRL').

False
color_key list, optional (default: None)

Custom colors for condition group labels. If None, uses default palette.

None
group_cond bool, optional (default: False)

If True and cond_group_labels provided, sorts conditions by group.

False
control_pattern str, optional (default: None)

Substring identifying the control conditions whose median and spread set the normalization. Overrides ThomsonNorm's implied 'CTRL'.

None
control_conditions Sequence[str], optional (default: None)

Explicit control condition names, taking precedence over control_pattern. With neither given, every condition contributes.

None
Source code in scrise/plotting/factors.py
def plot_condition_factors(
    data: anndata.AnnData,
    ax: Axes,
    cond: str = "Condition",
    log_transform: bool = True,
    cond_group_labels: pd.Series | None = None,
    ThomsonNorm: bool = False,
    color_key=None,
    group_cond: bool = False,
    control_pattern: str | None = None,
    control_conditions: Sequence[str] | None = None,
):
    """Plot condition factors as a heatmap showing how conditions contribute to
    components.

    This visualization shows how each experimental condition (rows) contributes to
    each RISE component (columns). High values indicate strong association between
    a condition and a component's pattern. Log transformation and normalization
    help reveal relative differences across conditions.

    Parameters
    ----------
    data : anndata.AnnData
        AnnData object with RISE decomposition results. Must contain:
        - data.uns["Pf2_A"]: Condition factors (n_conditions, rank)
        - data.obs[cond]: Condition labels for each cell
    ax : matplotlib.axes.Axes
        Matplotlib axes object to plot on.
    cond : str, optional (default: "Condition")
        Name of column in data.obs containing condition labels.
    log_transform : bool, optional (default: True)
        If True, applies log10 transformation to condition factors before plotting.
        This helps visualize differences when values span orders of magnitude.
    cond_group_labels : pandas.Series, optional (default: None)
        Series mapping conditions to group labels for colored row annotations.
        Useful for grouping related conditions (e.g., drug classes, patient cohorts).
    ThomsonNorm : bool, optional (default: False)
        If True, normalizes factors using only control conditions (those
        containing 'CTRL').
    color_key : list, optional (default: None)
        Custom colors for condition group labels. If None, uses default palette.
    group_cond : bool, optional (default: False)
        If True and cond_group_labels provided, sorts conditions by group.
    control_pattern : str, optional (default: None)
        Substring identifying the control conditions whose median and spread
        set the normalization. Overrides ``ThomsonNorm``'s implied 'CTRL'.
    control_conditions : Sequence[str], optional (default: None)
        Explicit control condition names, taking precedence over
        ``control_pattern``. With neither given, every condition contributes.
    """
    yt = pd.Series(np.unique(data.obs[cond]))
    X = np.array(data.uns["Pf2_A"])

    if log_transform is True:
        X = np.log10(X)

    XX = _normalization_rows(yt, X, ThomsonNorm, control_pattern, control_conditions)

    X -= np.median(XX, axis=0)
    X /= np.std(XX, axis=0)

    if log_transform is False:
        X -= np.min(X, axis=0)

    ind = reorder_table(X)
    X = X[ind]
    yt = yt.iloc[ind]

    if cond_group_labels is not None:
        cond_group_labels = cond_group_labels.iloc[ind]
        if group_cond is True:
            ind = cond_group_labels.argsort()
            cond_group_labels = cond_group_labels.iloc[ind]
            X = X[ind]
            yt = yt.iloc[ind]
        _draw_condition_group_labels(ax, cond_group_labels, color_key)

    xticks = np.arange(1, X.shape[1] + 1)
    sns.heatmap(
        data=X,
        xticklabels=xticks,
        yticklabels=yt,
        ax=ax,
        center=0,
        cmap=cmap,
    )
    ax.tick_params(axis="y", rotation=0)
    ax.set(xlabel="Component")

plot_eigenstate_factors

plot_eigenstate_factors(data: AnnData, ax: Axes)

Plot eigen-state factors as a heatmap showing cell state patterns.

Eigen-state factors represent the underlying cell state patterns across components. Each row represents an eigen-state (a summary of similar cells), and each column represents a component. High values indicate strong association between a cell state pattern and a component.

Parameters:

Name Type Description Default
data AnnData

AnnData object with RISE decomposition results. Must contain: - data.uns["Pf2_B"]: Eigen-state factors (rank, rank)

required
ax Axes

Matplotlib axes object to plot on.

required
Source code in scrise/plotting/factors.py
def plot_eigenstate_factors(data: anndata.AnnData, ax: Axes):
    """Plot eigen-state factors as a heatmap showing cell state patterns.

    Eigen-state factors represent the underlying cell state patterns across components.
    Each row represents an eigen-state (a summary of similar cells), and each column
    represents a component. High values indicate strong association between a cell
    state pattern and a component.

    Parameters
    ----------
    data : anndata.AnnData
        AnnData object with RISE decomposition results. Must contain:
        - data.uns["Pf2_B"]: Eigen-state factors (rank, rank)
    ax : matplotlib.axes.Axes
        Matplotlib axes object to plot on.
    """
    X = np.asarray(data.uns["Pf2_B"])
    _draw_unit_scaled_heatmap(X, np.arange(1, X.shape[1] + 1), ax)

plot_gene_factors

plot_gene_factors(
    data: AnnData, ax: Axes, weight=0.08, trim=True
)

Plot gene factors as a heatmap showing which genes contribute to each component.

This visualization reveals coordinated gene modules by showing which genes (rows) are highly weighted in each component (columns). The weight parameter filters out genes with low contributions, focusing on the most important genes for interpretation.

Parameters:

Name Type Description Default
data AnnData

AnnData object with RISE decomposition results. Must contain: - data.varm["Pf2_C"]: Gene factors (n_genes, rank)

required
ax Axes

Matplotlib axes object to plot on.

required
weight float, optional (default: 0.08)

Minimum absolute weight threshold for including genes. Genes with maximum absolute weight below this value across all components are filtered out. Higher values show fewer, more important genes.

0.08
trim bool, optional (default: True)

If True, filters genes based on the weight parameter. If False, shows all genes.

True
Source code in scrise/plotting/factors.py
def plot_gene_factors(data: anndata.AnnData, ax: Axes, weight=0.08, trim=True):
    """Plot gene factors as a heatmap showing which genes contribute to each component.

    This visualization reveals coordinated gene modules by showing which genes (rows)
    are highly weighted in each component (columns). The weight parameter filters out
    genes with low contributions, focusing on the most important genes for
    interpretation.

    Parameters
    ----------
    data : anndata.AnnData
        AnnData object with RISE decomposition results. Must contain:
        - data.varm["Pf2_C"]: Gene factors (n_genes, rank)
    ax : matplotlib.axes.Axes
        Matplotlib axes object to plot on.
    weight : float, optional (default: 0.08)
        Minimum absolute weight threshold for including genes. Genes with maximum
        absolute weight below this value across all components are filtered out.
        Higher values show fewer, more important genes.
    trim : bool, optional (default: True)
        If True, filters genes based on the weight parameter. If False, shows all genes.
    """
    X = np.array(data.varm["Pf2_C"])
    yt = data.var.index.values

    if trim is True:
        max_weight = np.max(np.abs(X), axis=1)
        kept_idxs = max_weight > weight
        if not np.any(kept_idxs):
            raise ValueError(
                f"No genes exceeded the weight threshold {weight}. Lower the threshold or set trim=False."
            )
        X = X[kept_idxs]
        yt = yt[kept_idxs]

    ind = reorder_table(X)
    _draw_unit_scaled_heatmap(X[ind], yt[ind], ax)

PaCMAP Visualization

scrise.plotting.pacmap

plot_labels_pacmap

plot_labels_pacmap(
    X: AnnData,
    labelType: str,
    ax: Axes,
    condition=None,
    cmap: str = "tab20",
    color_key=None,
)

Plot PaCMAP embedding colored by categorical labels (cell type or condition).

This visualization shows the overall structure of the cell embedding, revealing how cells cluster by cell type, experimental condition, or other categorical metadata. Useful for understanding the biological organization captured by RISE.

Parameters:

Name Type Description Default
X AnnData

AnnData object with RISE decomposition results. Must contain: - X.obsm["X_pf2_PaCMAP"]: PaCMAP embedding coordinates (n_cells, 2) - X.obs[labelType]: Categorical labels for coloring cells

required
labelType str

Name of column in X.obs containing categorical labels to color by. Common values: "Cell Type", "Condition", "Sample", etc.

required
ax Axes

Matplotlib axes object to plot on.

required
condition list of str, optional (default: None)

If provided, only highlights cells from these specific conditions/labels. All other cells are labeled as "Other".

None
cmap str, optional (default: "tab20")

Matplotlib colormap name for coloring categories.

'tab20'
color_key list, optional (default: None)

Custom list of colors for categories. If None, uses cmap.

None
Source code in scrise/plotting/pacmap.py
def plot_labels_pacmap(
    X: anndata.AnnData,
    labelType: str,
    ax: Axes,
    condition=None,
    cmap: str = "tab20",
    color_key=None,
):
    """Plot PaCMAP embedding colored by categorical labels (cell type or condition).

    This visualization shows the overall structure of the cell embedding, revealing
    how cells cluster by cell type, experimental condition, or other categorical
    metadata. Useful for understanding the biological organization captured by RISE.

    Parameters
    ----------
    X : anndata.AnnData
        AnnData object with RISE decomposition results. Must contain:
        - X.obsm["X_pf2_PaCMAP"]: PaCMAP embedding coordinates (n_cells, 2)
        - X.obs[labelType]: Categorical labels for coloring cells
    labelType : str
        Name of column in X.obs containing categorical labels to color by.
        Common values: "Cell Type", "Condition", "Sample", etc.
    ax : matplotlib.axes.Axes
        Matplotlib axes object to plot on.
    condition : list of str, optional (default: None)
        If provided, only highlights cells from these specific conditions/labels.
        All other cells are labeled as "Other".
    cmap : str, optional (default: "tab20")
        Matplotlib colormap name for coloring categories.
    color_key : list, optional (default: None)
        Custom list of colors for categories. If None, uses cmap.
    """
    labels = X.obs[labelType]

    if condition is not None:
        labels = pd.Series([c if c in condition else "Other" for c in labels])
    if labels.dtype == "category":
        labels = labels.cat.set_categories(
            np.sort(labels.cat.categories.values), ordered=True
        )
    indices = np.argsort(labels)

    points = X.obsm["X_pf2_PaCMAP"][indices, :]
    labels = labels.iloc[indices]

    data = pd.DataFrame(points, columns=["x", "y"])
    data["label"] = pd.Categorical(labels)

    unique_labels = np.unique(labels)
    num_labels = unique_labels.shape[0]
    if color_key is None:
        color_key = _to_hex(plt.get_cmap(cmap)(np.linspace(0, 1, num_labels)))
    legend_elements = [
        Patch(facecolor=color_key[i], label=k) for i, k in enumerate(unique_labels)
    ]

    _render_pacmap(
        points=points,
        data=data,
        agg_expr=ds.count_cat("label"),
        ax=ax,
        color_key=color_key,
        how="eq_hist",
    )
    ax.legend(handles=legend_elements)

plot_gene_pacmap

plot_gene_pacmap(
    gene: str, X: AnnData, ax: Axes, clip_outliers=0.9995
)

Plot PaCMAP embedding colored by gene expression levels.

This visualization overlays gene expression onto the PaCMAP embedding of cells, revealing which cell populations express specific genes. Useful for validating component interpretations by checking if marker genes align with component patterns.

Parameters:

Name Type Description Default
gene str

Name of gene to visualize. Must be present in X.var_names.

required
X AnnData

AnnData object with RISE decomposition results. Must contain: - X.obsm["X_pf2_PaCMAP"]: PaCMAP embedding coordinates (n_cells, 2) - X[:, gene]: Gene expression values - X.var["means"]: Pre-computed gene means for centering

required
ax Axes

Matplotlib axes object to plot on.

required
clip_outliers float, optional (default: 0.9995)

Quantile threshold for clipping extreme expression values. Values above this quantile are clipped to improve visualization contrast.

0.9995
Source code in scrise/plotting/pacmap.py
def plot_gene_pacmap(gene: str, X: anndata.AnnData, ax: Axes, clip_outliers=0.9995):
    """Plot PaCMAP embedding colored by gene expression levels.

    This visualization overlays gene expression onto the PaCMAP embedding of cells,
    revealing which cell populations express specific genes. Useful for validating
    component interpretations by checking if marker genes align with component patterns.

    Parameters
    ----------
    gene : str
        Name of gene to visualize. Must be present in X.var_names.
    X : anndata.AnnData
        AnnData object with RISE decomposition results. Must contain:
        - X.obsm["X_pf2_PaCMAP"]: PaCMAP embedding coordinates (n_cells, 2)
        - X[:, gene]: Gene expression values
        - X.var["means"]: Pre-computed gene means for centering
    ax : matplotlib.axes.Axes
        Matplotlib axes object to plot on.
    clip_outliers : float, optional (default: 0.9995)
        Quantile threshold for clipping extreme expression values.
        Values above this quantile are clipped to improve visualization contrast.
    """
    geneList = X[:, gene].to_df().values
    geneList = np.clip(geneList, None, np.quantile(geneList, clip_outliers))
    cmap = sns.color_palette("ch:s=-.2,r=.6", as_cmap=True)

    values = geneList - np.min(geneList)
    values /= np.max(values)

    points = np.array(X.obsm["X_pf2_PaCMAP"])
    _plot_continuous_pacmap(points, values, ax, cmap, (0.0, 1.0), f"{gene}")

plot_wp_pacmap

plot_wp_pacmap(
    X: AnnData, cmp: int, ax: Axes, cbarMax: float = 1.0
)

Plot PaCMAP embedding colored by weighted projections for a component.

This visualization shows which cells contribute most strongly to a specific component by coloring them according to their weighted projections. Cells with high weighted projections (bright colors) are most representative of that component's expression pattern.

Parameters:

Name Type Description Default
X AnnData

AnnData object with RISE decomposition results. Must contain: - X.obsm["X_pf2_PaCMAP"]: PaCMAP embedding coordinates (n_cells, 2) - X.obsm["weighted_projections"]: Weighted cell projections (n_cells, rank)

required
cmp int

Component number to visualize (1-indexed). For example, cmp=10 shows the cell associations for component 10.

required
ax Axes

Matplotlib axes object to plot on.

required
cbarMax float, optional (default: 1.0)

Maximum value for the color scale. Values are normalized to [-cbarMax, cbarMax]. Lower values increase contrast for components with weaker associations.

1.0
Source code in scrise/plotting/pacmap.py
def plot_wp_pacmap(X: anndata.AnnData, cmp: int, ax: Axes, cbarMax: float = 1.0):
    """Plot PaCMAP embedding colored by weighted projections for a component.

    This visualization shows which cells contribute most strongly to a specific
    component by coloring them according to their weighted projections. Cells with
    high weighted projections (bright colors) are most representative of that
    component's expression pattern.

    Parameters
    ----------
    X : anndata.AnnData
        AnnData object with RISE decomposition results. Must contain:
        - X.obsm["X_pf2_PaCMAP"]: PaCMAP embedding coordinates (n_cells, 2)
        - X.obsm["weighted_projections"]: Weighted cell projections (n_cells, rank)
    cmp : int
        Component number to visualize (1-indexed). For example, cmp=10 shows
        the cell associations for component 10.
    ax : matplotlib.axes.Axes
        Matplotlib axes object to plot on.
    cbarMax : float, optional (default: 1.0)
        Maximum value for the color scale. Values are normalized to [-cbarMax, cbarMax].
        Lower values increase contrast for components with weaker associations.
    """
    values = X.obsm["weighted_projections"][:, cmp - 1]
    points = np.asarray(X.obsm["X_pf2_PaCMAP"])
    cmap = sns.diverging_palette(240, 10, as_cmap=True)

    values /= np.max(np.abs(values))
    _plot_continuous_pacmap(
        points, values, ax, cmap, (-cbarMax, cbarMax), f"Cmp. {cmp}"
    )

Rank Selection Plotting

scrise.plotting.rank_selection

Plotting functions for BiCV-based rank selection.

plot_bicv_r2x

plot_bicv_r2x(results: DataFrame, ax: Axes) -> None

Plot BiCV R2X and in-sample fit R2X across ranks.

The fit R2X (in-sample, computed on the full dataset) increases monotonically with rank. The BiCV R2X (held-out, averaged across repeated random cell/gene splits) penalizes overfitting and typically peaks near the rank that best generalizes to unseen data. The peak (or plateau) of the BiCV R2X curve is a good candidate for the rank to use.

Parameters:

Name Type Description Default
results DataFrame

Output of :func:scrise.rank_selection.bicv, with columns "Rank", "Repeat", "Metric" ("Fit R2X" or "BiCV R2X"), and "R2X".

required
ax Axes

Matplotlib axes object to plot on.

required
Source code in scrise/plotting/rank_selection.py
def plot_bicv_r2x(results: pd.DataFrame, ax: Axes) -> None:
    """Plot BiCV R2X and in-sample fit R2X across ranks.

    The fit R2X (in-sample, computed on the full dataset) increases
    monotonically with rank. The BiCV R2X (held-out, averaged across
    repeated random cell/gene splits) penalizes overfitting and typically
    peaks near the rank that best generalizes to unseen data. The peak (or
    plateau) of the BiCV R2X curve is a good candidate for the rank to use.

    Parameters
    ----------
    results : pandas.DataFrame
        Output of :func:`scrise.rank_selection.bicv`, with columns "Rank",
        "Repeat", "Metric" ("Fit R2X" or "BiCV R2X"), and "R2X".
    ax : matplotlib.axes.Axes
        Matplotlib axes object to plot on.
    """
    sns.lineplot(
        data=results,
        x="Rank",
        y="R2X",
        hue="Metric",
        style="Metric",
        markers=True,
        dashes=False,
        errorbar="sd",
        ax=ax,
    )
    ax.set(xlabel="Rank", ylabel="R2X")
    ax.legend(title=None)

Factor Stability

scrise.plotting.stability

Decomposition stability analysis and plotting functions.

plot_fms_diff_ranks

plot_fms_diff_ranks(
    X: AnnData,
    ax: Axes,
    ranksList: list[int],
    runs: int,
    compress: int
    | tuple[int, int | None]
    | str
    | bool
    | None = "auto",
)

Plot Factor Match Score (FMS) across different ranks to assess stability.

FMS measures the reproducibility of PARAFAC2 decomposition results across multiple runs. Values above ~0.6 indicate stable, reproducible components. This helps determine which ranks produce reliable decompositions that are not overly sensitive to initialization or noise.

Parameters:

Name Type Description Default
X AnnData

Preprocessed AnnData object containing single-cell RNA-seq data. Must have X.obs["condition_unique_idxs"] for RISE decomposition.

required
ax Axes

Matplotlib axes object to plot on.

required
ranksList list of int

List of rank values to test (e.g., [1, 5, 10, 15, 20, 25, 30]). Each rank will be run multiple times to compute FMS.

required
runs int

Number of independent runs per rank to use for FMS calculation. Higher values give more reliable FMS estimates but take longer. Typical values: 3-5 runs.

required
compress int | tuple[int, int | None] | str | bool | None

CANDELINC compression mode passed to pf2 for each rank's fit. Defaults to "auto", which sharply cuts the cost of sweeping many ranks and runs over raw data. Pass None/False to fall back to exact ALS.

'auto'
Notes

FMS values interpretation: - FMS > 0.9: Highly stable decomposition - FMS > 0.6: Acceptably stable decomposition - FMS < 0.6: Unstable, consider lower rank or more data

Source code in scrise/plotting/stability.py
def plot_fms_diff_ranks(
    X: anndata.AnnData,
    ax: Axes,
    ranksList: list[int],
    runs: int,
    compress: int | tuple[int, int | None] | str | bool | None = "auto",
):
    """Plot Factor Match Score (FMS) across different ranks to assess stability.

    FMS measures the reproducibility of PARAFAC2 decomposition results across
    multiple runs. Values above ~0.6 indicate stable, reproducible components.
    This helps determine which ranks produce reliable decompositions that are
    not overly sensitive to initialization or noise.

    Parameters
    ----------
    X : anndata.AnnData
        Preprocessed AnnData object containing single-cell RNA-seq data.
        Must have X.obs["condition_unique_idxs"] for RISE decomposition.
    ax : matplotlib.axes.Axes
        Matplotlib axes object to plot on.
    ranksList : list of int
        List of rank values to test (e.g., [1, 5, 10, 15, 20, 25, 30]).
        Each rank will be run multiple times to compute FMS.
    runs : int
        Number of independent runs per rank to use for FMS calculation.
        Higher values give more reliable FMS estimates but take longer.
        Typical values: 3-5 runs.
    compress : int | tuple[int, int | None] | str | bool | None, optional
        CANDELINC compression mode passed to ``pf2`` for each rank's fit.
        Defaults to ``"auto"``, which sharply cuts the cost of sweeping many
        ranks and runs over raw data. Pass None/False to fall back to exact
        ALS.

    Notes
    -----
    FMS values interpretation:
    - FMS > 0.9: Highly stable decomposition
    - FMS > 0.6: Acceptably stable decomposition
    - FMS < 0.6: Unstable, consider lower rank or more data
    """
    fmsLists = []
    for j in range(runs):
        scores = []
        for i in ranksList:
            dataX = pf2(X, rank=i, random_state=j, doEmbedding=False, compress=compress)
            sampledX = pf2(
                resample(X),
                rank=i,
                random_state=j,
                doEmbedding=False,
                compress=compress,
            )

            fmsScore = calculateFMS(dataX, sampledX)
            scores.append(fmsScore)
        fmsLists.append(scores)

    df = pd.DataFrame(fmsLists, columns=ranksList).melt(
        var_name="Component", value_name="FMS"
    )
    df["Run"] = np.repeat(np.arange(runs), len(ranksList))

    sns.lineplot(data=df, x="Component", y="FMS", ax=ax)
    ax.set_ylim(0, 1)

Cell-Type Alignment Plotting

scrise.plotting.annotation_alignment

Plotting functions for cell-type alignment scoring.

plot_cell_type_alignment

plot_cell_type_alignment(
    data: AnnData | CellTypeAlignmentResults | DataFrame,
    ax: Axes | Sequence[Axes] | None = None,
    cell_type_col: str = "cell_type",
    projection_key: str = "weighted_projections",
    signed: bool = False,
    n_permutations: int = 1000,
    alpha: float = 0.05,
    annotate_significance: bool = True,
    show_metrics: bool = True,
    reorder: bool = True,
    cmap=None,
    metrics_cmap="Blues",
    random_state=None,
) -> Axes | tuple[Axes, ...]

Plot cell-type alignment heatmap for RISE components.

Visualizes per-cell-type AUROC enrichment for each component with optional significance stars (* for q <= alpha) and row annotations for uniqueness (tau) and variance explained (eta^2).

Parameters:

Name Type Description Default
data AnnData | CellTypeAlignmentResults | DataFrame

Alignment results or AnnData object containing RISE results.

required
ax Axes | Sequence[Axes] | None

Matplotlib Axes to plot on. Can be: - None: creates a new figure with appropriate subplots. - Single Axes: plots the main AUROC heatmap on the provided axes. - Pair of Axes (ax_main, ax_metrics): plots heatmap on ax_main and metrics on ax_metrics.

None
cell_type_col str, optional (default: "cell_type")

Column name in data.obs containing cell type labels (if data is AnnData).

'cell_type'
projection_key str, optional (default: "weighted_projections")

Key in data.obsm containing projections (if data is AnnData).

'weighted_projections'
signed bool, optional (default: False)

If True, takes the absolute value of loadings.

False
n_permutations int, optional (default: 1000)

Number of permutations to run (if scoring an AnnData).

1000
alpha float, optional (default: 0.05)

FDR significance threshold.

0.05
annotate_significance bool, optional (default: True)

If True, marks significant cell types (q <= alpha, AUROC > 0.5) with an asterisk (*).

True
show_metrics bool, optional (default: True)

If True and subplots are available (or ax is None), shows tau and eta^2 row annotations.

True
reorder bool, optional (default: True)

If True, reorders components by dominant cell type.

True
cmap Colormap or str

Colormap for AUROC enrichment heatmap (defaults to blue-red diverging).

None
metrics_cmap Colormap or str, optional (default: "Blues")

Colormap for tau and eta^2 annotations.

'Blues'
random_state int | None

Random seed for permutations.

None

Returns:

Type Description
Axes | tuple[Axes, ...]

The plotted Matplotlib Axes.

Source code in scrise/plotting/annotation_alignment.py
def plot_cell_type_alignment(
    data: anndata.AnnData | CellTypeAlignmentResults | pd.DataFrame,
    ax: Axes | Sequence[Axes] | None = None,
    cell_type_col: str = "cell_type",
    projection_key: str = "weighted_projections",
    signed: bool = False,
    n_permutations: int = 1000,
    alpha: float = 0.05,
    annotate_significance: bool = True,
    show_metrics: bool = True,
    reorder: bool = True,
    cmap=None,
    metrics_cmap="Blues",
    random_state=None,
) -> Axes | tuple[Axes, ...]:
    """Plot cell-type alignment heatmap for RISE components.

    Visualizes per-cell-type AUROC enrichment for each component with optional
    significance stars (* for q <= alpha) and row annotations for uniqueness (tau)
    and variance explained (eta^2).

    Parameters
    ----------
    data : anndata.AnnData | CellTypeAlignmentResults | pd.DataFrame
        Alignment results or AnnData object containing RISE results.
    ax : Axes | Sequence[Axes] | None, optional
        Matplotlib Axes to plot on. Can be:
        - None: creates a new figure with appropriate subplots.
        - Single Axes: plots the main AUROC heatmap on the provided axes.
        - Pair of Axes (ax_main, ax_metrics): plots heatmap on ax_main and metrics on ax_metrics.
    cell_type_col : str, optional (default: "cell_type")
        Column name in data.obs containing cell type labels (if data is AnnData).
    projection_key : str, optional (default: "weighted_projections")
        Key in data.obsm containing projections (if data is AnnData).
    signed : bool, optional (default: False)
        If True, takes the absolute value of loadings.
    n_permutations : int, optional (default: 1000)
        Number of permutations to run (if scoring an AnnData).
    alpha : float, optional (default: 0.05)
        FDR significance threshold.
    annotate_significance : bool, optional (default: True)
        If True, marks significant cell types (q <= alpha, AUROC > 0.5) with an asterisk (*).
    show_metrics : bool, optional (default: True)
        If True and subplots are available (or ax is None), shows tau and eta^2 row annotations.
    reorder : bool, optional (default: True)
        If True, reorders components by dominant cell type.
    cmap : Colormap or str, optional
        Colormap for AUROC enrichment heatmap (defaults to blue-red diverging).
    metrics_cmap : Colormap or str, optional (default: "Blues")
        Colormap for tau and eta^2 annotations.
    random_state : int | None, optional
        Random seed for permutations.

    Returns
    -------
    Axes | tuple[Axes, ...]
        The plotted Matplotlib Axes.
    """
    results = _coerce_alignment_results(
        data,
        cell_type_col,
        projection_key,
        signed,
        n_permutations,
        alpha,
        random_state,
    )

    enrichment_df = results.enrichment.copy()
    q_values_df = results.q_values.copy()
    tau_series = results.tau.copy()
    eta2_series = results.eta_squared.copy()

    if reorder and len(enrichment_df) > 1:
        enrichment_df, q_values_df, tau_series, eta2_series = (
            _order_by_dominant_cell_type(
                enrichment_df, q_values_df, tau_series, eta2_series
            )
        )

    annot_df = None
    if annotate_significance:
        annot_df = _significance_annotations(enrichment_df, q_values_df, results.alpha)

    heatmap_cmap = cmap if cmap is not None else cmap_enrichment
    ax_main, ax_metrics = _resolve_alignment_axes(ax, show_metrics, enrichment_df)

    # Main AUROC heatmap
    sns.heatmap(
        data=enrichment_df,
        ax=ax_main,
        cmap=heatmap_cmap,
        center=0.5,
        vmin=0.0,
        vmax=1.0,
        cbar_kws={"label": "AUROC (Enrichment)"},
        annot=annot_df if annotate_significance else False,
        fmt="",
        annot_kws={"size": 14, "va": "center", "ha": "center", "weight": "bold"},
    )
    ax_main.set_ylabel("Component")
    ax_main.set_xlabel("Cell Type")
    ax_main.tick_params(axis="y", rotation=0)
    ax_main.tick_params(axis="x", rotation=45)

    # Side metrics heatmap
    if ax_metrics is not None and show_metrics:
        metrics_df = pd.DataFrame(
            {
                r"$\tau$": tau_series,
                r"$\eta^2$": eta2_series,
            },
            index=enrichment_df.index,
        )
        sns.heatmap(
            data=metrics_df,
            ax=ax_metrics,
            cmap=metrics_cmap,
            vmin=0.0,
            vmax=1.0,
            cbar_kws={"label": "Score"},
            annot=True,
            fmt=".2f",
            annot_kws={"size": 9},
            yticklabels=False,
        )
        ax_metrics.set_ylabel("")
        ax_metrics.tick_params(axis="x", rotation=0)

    if ax_metrics is not None:
        return (ax_main, ax_metrics)
    return ax_main