Source code for pyglmGamPoi.validate

"""Input validation and output normalization."""

from __future__ import annotations

from typing import Union

import numpy as np
import pandas as pd
from scipy import sparse

from pyglmGamPoi.constants import LOG10_SLOPE, LOG_UMI_COLUMN, MODEL_PARS_COLUMNS


[docs] def normalize_model_pars_columns(model_pars: pd.DataFrame) -> pd.DataFrame: """ Normalize coefficient column names to trackcell conventions. trackcell applies the same rename in ``_sctransform_v2._normalize_model_pars_columns``. """ out = model_pars.copy() rename = {col: col.strip("()") for col in out.columns if col.startswith("(")} if rename: out = out.rename(columns=rename) if "Intercept" not in out.columns and "(Intercept)" in out.columns: out = out.rename(columns={"(Intercept)": "Intercept"}) return out
def validate_genes_by_cells_matrix( umi: Union[np.ndarray, sparse.spmatrix], *, name: str = "umi", ) -> tuple[sparse.csr_matrix, int, int]: """ Ensure a genes × cells count matrix (trackcell VST internal layout). AnnData stores cells × genes; callers must transpose before calling pyglmGamPoi unless using :func:`pyglmGamPoi.adata.extract_genes_by_cells_from_adata`. """ if sparse.issparse(umi): mat = umi.tocsr().astype(np.float64) else: arr = np.asarray(umi, dtype=np.float64) if arr.ndim != 2: raise ValueError(f"{name} must be 2-D (genes × cells)") mat = sparse.csr_matrix(arr) if mat.shape[0] == 0 or mat.shape[1] == 0: raise ValueError(f"{name} must be non-empty") if (mat.data < 0).any(): raise ValueError(f"{name} must contain non-negative counts") return mat, mat.shape[0], mat.shape[1] def validate_regressor_data( regressor_data: pd.DataFrame, n_cells: int, *, cell_index: pd.Index | None = None, ) -> pd.DataFrame: """ Validate cell-level covariates for offset fitting. ``log_umi`` must be in **log10** scale (trackcell ``cell_attr`` convention). """ if not isinstance(regressor_data, pd.DataFrame): raise TypeError("regressor_data must be a pandas DataFrame") if LOG_UMI_COLUMN not in regressor_data.columns: raise ValueError( f"regressor_data must contain column '{LOG_UMI_COLUMN}' (log10 UMI per cell)" ) if len(regressor_data) != n_cells: raise ValueError( f"regressor_data has {len(regressor_data)} rows but umi has {n_cells} cells" ) if cell_index is not None and not regressor_data.index.equals(cell_index): raise ValueError( "regressor_data.index must match the cell index of umi " "(same labels and order)" ) log10_umi = regressor_data[LOG_UMI_COLUMN].to_numpy(dtype=np.float64) if not np.isfinite(log10_umi).all(): raise ValueError(f"regressor_data['{LOG_UMI_COLUMN}'] contains non-finite values") return regressor_data def validate_gene_index(gene_index: pd.Index, n_genes: int) -> pd.Index: if not isinstance(gene_index, pd.Index): gene_index = pd.Index(gene_index) if len(gene_index) != n_genes: raise ValueError( f"gene_index length ({len(gene_index)}) != number of genes ({n_genes})" ) return gene_index.astype(str) def format_model_pars( theta: np.ndarray, intercept: np.ndarray, gene_index: pd.Index, *, extra_columns: dict[str, np.ndarray] | None = None, ) -> pd.DataFrame: """Build the canonical model_pars DataFrame expected by trackcell VST.""" if len(theta) != len(gene_index) or len(intercept) != len(gene_index): raise ValueError("theta/intercept length must match gene_index") data: dict[str, np.ndarray] = { "theta": np.asarray(theta, dtype=np.float64), "Intercept": np.asarray(intercept, dtype=np.float64), "log_umi": np.full(len(gene_index), LOG10_SLOPE, dtype=np.float64), } if extra_columns: for key, values in extra_columns.items(): if len(values) != len(gene_index): raise ValueError(f"extra column '{key}' length mismatch") data[key] = np.asarray(values, dtype=np.float64) out = pd.DataFrame(data, index=gene_index) out.index.name = None ordered = list(MODEL_PARS_COLUMNS) + [c for c in out.columns if c not in MODEL_PARS_COLUMNS] return out[ordered] def assert_model_pars_contract(model_pars: pd.DataFrame) -> pd.DataFrame: """Verify return value satisfies trackcell ``fit_offset_model`` contract.""" out = normalize_model_pars_columns(model_pars) missing = [c for c in MODEL_PARS_COLUMNS if c not in out.columns] if missing: raise ValueError(f"model_pars missing required columns: {missing}") if out["log_umi"].nunique() != 1 or not np.isclose(out["log_umi"].iloc[0], LOG10_SLOPE): raise ValueError("model_pars['log_umi'] must be constant log(10)") if not pd.api.types.is_string_dtype(out.index): out.index = out.index.astype(str) return out