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