diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..9b12267 --- /dev/null +++ b/.gitignore @@ -0,0 +1,34 @@ +# Byte-compiled / cache +__pycache__/ +*.py[cod] +*$py.class + +# Build / packaging artifacts +build/ +dist/ +*.egg-info/ +.eggs/ +*.egg + +# Test / coverage +.pytest_cache/ +.coverage +.coverage.* +htmlcov/ + +# Type/lint caches +.mypy_cache/ +.ruff_cache/ + +# Jupyter +.ipynb_checkpoints/ + +# Virtual environments +venv/ +env/ +.venv/ + +# Editors / OS +.vscode/ +.idea/ +.DS_Store diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 0000000..c7864a4 --- /dev/null +++ b/CHANGELOG.md @@ -0,0 +1,34 @@ +# Changelog + +## Breaking changes + +- `setup_reference()` no longer runs QC filtering or HVG selection automatically. It raises an error and tells you which `scanpy` function to run first. +- `setup_reference()` dropped the `copy` parameter. +- `setup_reference()` and `setup_spatial()` return `AnnData` instead of `dict`. +- `remove_mt` was renamed to `warn_mt` in both functions. MT/HLA/RPL genes now get a warning instead of automatic removal. +- `DOT.fit()` no longer accepts `mode`, `ratios_weight`, or `max_spot_size`. Pass these to `DOT()` when you construct it instead. `max_spot_size` was renamed to `max_size`. +- `DOT.get_weights()` returns a `pandas.DataFrame` instead of a numpy array. +- `kmeans_define_subtypes()` now returns a `DataFrame`. + +## New features + +- `setup_reference()` now supports predefined subtype annotations. This works alongside the existing k-means discovery. New functions: `predefined_subtypes()`, `check_subtype_consistency()`, `summarize_subtypes()`, `select_de_genes()`, `validate_reference_input()`. +- `DOT()` accepts `gene_weight`, `spot_weight`, `spatial_weight`, `ratios_weight`, `sparsity_coef`, and `cluster_weight` directly. You can also override any lambda value directly, or pass your own `weights_to_lambdas` function. +- `cluster_weight` implements `l_c`. This is the cluster-wise cosine term from the R package. It matches R's behaviour and is off by default. +- `setup_reference()` exposes the k-means subtype-count heuristic and per-type cell cap through `kws_kmeans`. Defaults match the previous hardcoded values. It also adds `min_frac` to filter out small clusters as noise. +- `setup_reference()` can further sub-cluster predefined subtypes that are under-resolved relative to the k-means heuristic. Set `refine_undersized=True` and pass `kws_refine` to control it. New functions: `plan_subtype_refinement()`, `refine_predefined_subtypes()`. +- `setup_reference()` stores a diagnostic table in `uns['subtype_plan']` when using predefined sub-clusters, with per-label cell counts, expected-vs-current cluster counts, and (when `refine_undersized=True`) how many sub-clusters were actually produced. +- `kmeans_define_subtypes()` exposes `max_genes` to control the gene-narrowing threshold, previously hardcoded at 500. +- `DOT.print_config()` prints the resolved optimisation configuration. +- `get_weights()` and `get_cell_types()` accept a `level` parameter. This aggregates results at any level of the subtype hierarchy. +- `plot_spatial_weights()` and `plot_cell_type_proportions()` accept a `DataFrame` directly. +- Added `dotpy.ds.toy_reference()` for generating synthetic test data. + +## Fixes + +- Fixed `_safe_log2`. It returned different values than R's `safelog2` for zero and negative inputs. This affects the abundance-matching and spatial-coherence loss terms. + +## Other + +- Added a `pytest` test suite with about 349 tests and 97% line coverage. +- Updated `run_dot_cli.py`, `example.py`, `README.md`, and `QUICKSTART.md` for the new API. diff --git a/QUICKSTART.md b/QUICKSTART.md index cf6c2fb..f4f5863 100644 --- a/QUICKSTART.md +++ b/QUICKSTART.md @@ -77,6 +77,14 @@ from dotpy import setup_reference, setup_spatial, DOT, plot_spatial_weights ref_adata = sc.read_h5ad('your_reference.h5ad') spatial_adata = sc.read_h5ad('your_spatial.h5ad') +# setup_reference() expects pre-filtered (positive) raw counts +# QC, MT/HLA/RPL removal, and HVG selection are standard scanpy steps left to you. +sc.pp.filter_cells(ref_adata, min_counts=1) +sc.pp.filter_genes(ref_adata, min_cells=1) +sc.pp.highly_variable_genes(adata, batch_key='sample') + +adata = adata[:, adata.var['highly_variable_intersection']] + # 2. Process data print("Processing reference...") ref_processed = setup_reference( @@ -102,17 +110,15 @@ print(f"Running on: {device}") dot = DOT( spatial_processed, ref_processed, + mode='highres', # For Xenium, MERFISH, CosMx, etc. + # mode='lowres', max_size=20, # OR for Visium, ST, etc. batch_size=500, # Adjust based on GPU memory device=device ) # 4. Run deconvolution print("Running DOT...") -# For high-resolution data (Xenium, MERFISH, CosMx) -dot.fit(mode='highres', iterations=100, verbose=True) - -# OR for low-resolution data (Visium, ST) -# dot.fit(mode='lowres', max_spot_size=20, iterations=100, verbose=True) +dot.fit(iterations=100, verbose=True) # 5. Get results weights = dot.get_weights(normalize=True) @@ -132,8 +138,8 @@ plot_spatial_weights( # 7. Save results spatial_adata.obsm['dot_weights'] = weights -for i, ct in enumerate(cell_types): - spatial_adata.obs[f'dot_{ct}'] = weights[:, i] +for ct in cell_types: + spatial_adata.obs[f'dot_{ct}'] = weights[ct] spatial_adata.write('spatial_deconvolved.h5ad') print("Results saved!") @@ -143,17 +149,26 @@ print("Results saved!") ### Reference Processing +`setup_reference()` only does sub-clustering + DE gene selection -- it expects +`adata` to already be QC-filtered (no empty cells/genes) with raw (positive) counts in +`.X`. It validates this instead of fixing it for you: it raises on empty +cells/genes, negative values, or non-integer counts (unless `check_counts=False`, +e.g. for background-corrected data), and on more than `max_input_genes` genes. +MT-/HLA-/RPL-prefixed genes are never removed automatically -- if present, +you'll get a warning with per-prefix counts (silence with `warn_mt=False`). + ```python ref_processed = setup_reference( adata, cell_type_key='cell_type', # Column with cell type labels subcluster_size=10, # Max subclusters per cell type (higher = more granular) - max_genes=5000, # Number of genes to use (higher = more info, slower) - remove_mt=True, # Remove mitochondrial genes + max_genes=5000, # Genes to keep in the final DE panel + max_input_genes=5000, # Genes adata is allowed to have on input (independent of max_genes) th_inner_logfold=0.75, # Log-fold threshold for gene selection in subclustering random_state=42, # Random seed for reproducibility - verbose=True, # Print progress - copy=True # Copy adata before processing + warn_mt=True, # Warn (don't remove) if MT/HLA/RPL genes are present + check_counts=True, # Verify adata.X looks like raw counts + verbose=True # Print progress ) ``` @@ -161,6 +176,10 @@ ref_processed = setup_reference( - Increase `subcluster_size` for more heterogeneous cell types - Increase `max_genes` if you have many similar cell types - Adjust `th_inner_logfold` to control gene selection stringency +- Raise `max_input_genes` if your reference has more genes than the default + 5000 and you'd rather not (re-)run your own HVG selection first +- Set `check_counts=False` if your counts aren't strictly integer (e.g. + ambient-RNA/background-corrected data) ### Spatial Processing @@ -171,7 +190,7 @@ spatial_processed = setup_spatial( th_spatial=0.84, # Similarity threshold for spatial neighbors th_gene_low=0.01, # Min expression frequency th_gene_high=0.99, # Max expression frequency - remove_mt=True, # Remove mitochondrial genes + warn_mt=True, # Warn if MT/HLA/RPL genes are present radius='auto', # Spatial neighborhood radius verbose=True, copy=True @@ -206,11 +225,17 @@ dot = DOT( ### DOT Fitting +`mode`, `ratios_weight`, and `max_size` are configured when constructing `DOT` +(see above), not at `fit()` time: + ```python # High-resolution (subcellular) +dot = DOT( + spatial_processed, ref_processed, + mode='highres', # For Xenium, MERFISH, CosMx, etc. + ratios_weight=0.0, # Weight for matching reference abundances (0-1) +) dot.fit( - mode='highres', # For Xenium, MERFISH, CosMx, etc. - ratios_weight=0.0, # Weight for matching reference abundances (0-1) iterations=100, # Number of optimization iterations gap_threshold=0.01, # Convergence threshold use_mixed_precision=False, # Use float16 on GPU (saves memory) @@ -221,10 +246,13 @@ dot.fit( ) # Low-resolution (spot-based) +dot = DOT( + spatial_processed, ref_processed, + mode='lowres', # For Visium, ST, etc. + ratios_weight=0.3, # Higher weight to match reference proportions + max_size=20, # Max cells per spot +) dot.fit( - mode='lowres', # For Visium, ST, etc. - ratios_weight=0.3, # Higher weight to match reference proportions - max_spot_size=20, # Max cells per spot iterations=100, gap_threshold=0.01, verbose=True @@ -260,7 +288,7 @@ spatial_adata.var_names = spatial_adata.var_names.str.upper() dot = DOT(spatial_processed, ref_processed, batch_size=100) # Option 2: Use mixed precision -dot.fit(mode='highres', use_mixed_precision=True) +dot.fit(use_mixed_precision=True) # Option 3: Reduce genes ref_processed = setup_reference( @@ -354,7 +382,6 @@ if torch.cuda.is_available(): ```python # Save checkpoints during long runs dot.fit( - mode='highres', iterations=200, checkpoint_dir='./checkpoints', checkpoint_freq=20, # Save every 20 iterations @@ -363,7 +390,6 @@ dot.fit( # Resume from checkpoint dot.fit( - mode='highres', iterations=300, # Continue to 300 total resume_from='./checkpoints/checkpoint_iter_200.pkl', verbose=True diff --git a/README.md b/README.md index ce77d9b..502fb5f 100644 --- a/README.md +++ b/README.md @@ -47,6 +47,13 @@ from dotpy import DOT, setup_reference, setup_spatial, plot_spatial_weights ref_adata = sc.read_h5ad('reference.h5ad') spatial_adata = sc.read_h5ad('spatial.h5ad') +# setup_reference() expects pre-filtered raw (positive) counts +# QC, MT/HLA/RPL removal, and HVG selection are standard scanpy steps left to you. +sc.pp.filter_cells(ref_adata, min_counts=1) +sc.pp.filter_genes(ref_adata, min_cells=1) +sc.pp.highly_variable_genes(adata, batch_key='sample') +adata = adata[:,adata.var['highly_variable_intersection']] + # Process reference and spatial data ref_processed = setup_reference( ref_adata, @@ -67,11 +74,11 @@ spatial_processed = setup_spatial( dot = DOT( spatial_processed, ref_processed, + mode='highres', batch_size=500 # Adjust for your GPU memory ) dot.fit( - mode='highres', iterations=100, checkpoint_dir='./checkpoints', # Save checkpoints checkpoint_freq=10, @@ -96,7 +103,6 @@ plot_spatial_weights( ```python dot.fit( - mode='highres', iterations=100, resume_from='./checkpoints/checkpoint_iter_50.pkl', verbose=True @@ -245,12 +251,13 @@ results/ For subcellular resolution data where each spot typically contains 1 cell: ```python -dot.fit( +dot = DOT( + spatial_processed, + ref_processed, mode='highres', ratios_weight=0.0, - iterations=100, - verbose=True ) +dot.fit(iterations=100, verbose=True) ``` ### Low-Resolution Data (Visium, ST) @@ -258,13 +265,14 @@ dot.fit( For spot-based technologies where spots contain multiple cells: ```python -dot.fit( +dot = DOT( + spatial_processed, + ref_processed, mode='lowres', - max_spot_size=20, # Maximum cells per spot + max_size=20, # Maximum cells per spot ratios_weight=0.3, # Weight for matching cell type proportions - iterations=100, - verbose=True ) +dot.fit(iterations=100, verbose=True) ``` ## Algorithm Overview @@ -284,14 +292,17 @@ The optimization is performed using the Frank-Wolfe algorithm, which efficiently ```python # Setup reference with custom parameters +# (ref_adata must already be QC-filtered, subset to highly variable genes, with raw counts in .X) ref_processed = setup_reference( ref_adata, cell_type_key='cell_type', subcluster_size=15, # More subclusters per cell type - max_genes=10000, # Use more genes - remove_mt=True, # Remove mitochondrial genes + max_genes=5000, # Genes to keep in the final DE panel + max_input_genes=5000, # Genes ref_adata is allowed to have on input th_inner_logfold=0.75, # Log-fold threshold for gene selection random_state=42, # For reproducibility + warn_mt=True, # Warn (don't remove) if MT/HLA/RPL genes are present + check_counts=True, # Verify adata.X looks like raw counts verbose=True ) @@ -303,7 +314,7 @@ spatial_processed = setup_spatial( th_gene_low=0.01, # Minimum gene expression frequency th_gene_high=0.99, # Maximum gene expression frequency radius='auto', # Or specify numeric value - remove_mt=True, # Remove mitochondrial genes + warn_mt=True, # Warn if MT/HLA/RPL genes are present verbose=True ) @@ -314,14 +325,14 @@ device = 'cuda' if torch.cuda.is_available() else 'cpu' dot = DOT( spatial_processed, ref_processed, + mode='highres', + ratios_weight=0.2, # Weight for abundance matching batch_size=500, # Adjust for GPU memory device=device # Explicitly set device ) # Fine-tune optimization dot.fit( - mode='highres', - ratios_weight=0.2, # Weight for abundance matching iterations=200, # More iterations gap_threshold=0.001, # Tighter convergence use_mixed_precision=True, # Use float16 on GPU @@ -348,12 +359,12 @@ dot = DOT( ### Saving Results ```python -# Add results to spatial AnnData -spatial_adata.obsm['dot_weights'] = weights +# Add results to spatial AnnData (weights is a DataFrame indexed by spot) +spatial_adata.obsm['dot_weights'] = weights.loc[spatial_adata.obs.index,:] # Add individual cell type columns -for i, ct in enumerate(cell_types): - spatial_adata.obs[f'dot_{ct}'] = weights[:, i] +for ct in cell_types: + spatial_adata.obs[f'dot_{ct}'] = weights.loc[spatial_adata.obs.index,ct] # Save spatial_adata.write('spatial_with_deconvolution.h5ad') @@ -436,11 +447,10 @@ ref_processed = setup_reference( ) # Use smaller batch size -dot = DOT(spatial, ref, batch_size=100) +dot = DOT(spatial, ref, mode='highres', batch_size=100) # Enable mixed precision on GPU dot.fit( - mode='highres', use_mixed_precision=True, iterations=100 ) @@ -450,11 +460,10 @@ dot.fit( ```python # Faster (fewer iterations) -dot.fit(mode='highres', iterations=50) +dot.fit(iterations=50) # More accurate (more iterations, tighter convergence) dot.fit( - mode='highres', iterations=200, gap_threshold=0.001 ) @@ -492,7 +501,7 @@ Nat Commun 15, 4994 (2024). https://doi.org/10.1038/s41467-024-48868-z dot = DOT(spatial, ref, batch_size=100) # Solution 2: Enable mixed precision -dot.fit(mode='highres', use_mixed_precision=True) +dot.fit(use_mixed_precision=True) # Solution 3: Use CPU dot = DOT(spatial, ref, device='cpu') diff --git a/dotpy/_exceptions.py b/dotpy/_exceptions.py new file mode 100644 index 0000000..8064d43 --- /dev/null +++ b/dotpy/_exceptions.py @@ -0,0 +1,38 @@ +"""Internal warning/exception utilities.""" + +import inspect +from pathlib import Path + +_PKG_DIR = Path(__file__).resolve().parent + + +def find_stack_level() -> int: + """ + Find the ``stacklevel`` that attributes a warning to the first caller + outside of dotpy, regardless of how many internal helpers it passed + through. + + Usage + ----- + ``warnings.warn(msg, UserWarning, stacklevel=find_stack_level())`` + """ + frame = inspect.currentframe() + n = 0 + while frame is not None: + frame = frame.f_back + n += 1 + if frame is None or not _is_in_package(frame): + break + return n + + +def _is_in_package(frame) -> bool: + try: + path = Path(frame.f_code.co_filename).resolve() + except OSError: + return False + try: + path.relative_to(_PKG_DIR) + except ValueError: + return False + return True diff --git a/dotpy/core.py b/dotpy/core.py index 508ab34..182ff3a 100644 --- a/dotpy/core.py +++ b/dotpy/core.py @@ -1,31 +1,298 @@ import numpy as np +import pandas as pd import torch import torch.nn.functional as F -from typing import Optional, Dict -from scipy.sparse import issparse +import warnings +from typing import Optional, Dict, Union, Callable +from anndata import AnnData +from scipy.sparse import issparse, triu from scipy import linalg as sp_linalg import time import pickle from pathlib import Path +from ._exceptions import find_stack_level +from .preprocessing import _symmetric_pairs_matrix + + +def _ref_to_anndata(ref: Union[Dict, AnnData]) -> AnnData: + """Accept the legacy dict shape or a real AnnData; normalize to AnnData.""" + if isinstance(ref, AnnData): + return ref + if not isinstance(ref, dict): + raise TypeError(f"ref must be a dict or AnnData, got {type(ref).__name__}") + + warnings.warn( + "Passing a dict for ref is deprecated, pass the AnnData returned by " + "setup_reference() instead.", + FutureWarning, + stacklevel=find_stack_level(), + ) + clusters = ref['clusters'] + major = np.empty(ref['X_sparse'].shape[0], dtype=object) + for ct, idx in clusters.items(): + major[np.asarray(idx)] = ct + obs = pd.DataFrame({'major': major}) + obs.index = obs.index.astype(str) + var = pd.DataFrame(index=np.asarray(ref['genes'])) + adata = AnnData(X=ref['X_sparse'], obs=obs, var=var) + adata.uns['ratios'] = ref['ratios'] + adata.uns['level_keys'] = ['major'] + return adata + + +def _spatial_to_anndata(spatial: Union[Dict, AnnData]) -> AnnData: + """Accept the legacy dict shape or a real AnnData; normalize to AnnData.""" + if isinstance(spatial, AnnData): + return spatial + if not isinstance(spatial, dict): + raise TypeError(f"spatial must be a dict or AnnData, got {type(spatial).__name__}") + + warnings.warn( + "Passing a dict for spatial is deprecated, pass the AnnData returned " + "by setup_spatial() instead.", + FutureWarning, + stacklevel=find_stack_level(), + ) + n_spots = spatial['X_sparse'].shape[0] + obs = pd.DataFrame(index=[str(i) for i in range(n_spots)]) + var = pd.DataFrame(index=np.asarray(spatial['genes'])) + adata = AnnData(X=spatial['X_sparse'], obs=obs, var=var) + adata.obsm['spatial'] = np.asarray(spatial['coords']) + + pairs = spatial.get('pairs') + if pairs is not None: + adata.obsp['dot_spatial_pairs'] = _symmetric_pairs_matrix( + np.asarray(pairs['i']), np.asarray(pairs['j']), np.asarray(pairs['w']), n=n_spots + ) + return adata + + +def _clusters_at_level(ref: AnnData, level: Optional[str] = None) -> Dict[str, np.ndarray]: + """ + Cell type at the given hierarchy level -> positional row indices in ref. + Defaults to the coarsest level; raises if level isn't a valid one. + """ + level_keys = ref.uns['level_keys'] + if level is None: + level = level_keys[0] + elif level not in level_keys: + raise ValueError(f"level must be one of {list(level_keys)}, got {level!r}") + + labels = ref.obs[level].values + return {ct: np.where(labels == ct)[0] for ct in pd.unique(labels)} + + +# Mode is a convenience preset: it only supplies *defaults* for the +# weights below, individually-passed values always win over it. spot_weight +# is deliberately absent here -- its default depends on the resolved +# max_size (see __init__), not on mode directly. +MODE_PRESETS: Dict[str, Dict[str, float]] = { + 'highres': dict(gene_weight=1.0, spatial_weight=0.01, ratios_weight=0.0, sparsity_coef=0.6, cluster_weight=0.0, max_size=1), + 'lowres': dict(gene_weight=1.0, spatial_weight=0.01, ratios_weight=0.0, sparsity_coef=0.4, cluster_weight=0.0, max_size=20), +} + +_LAMBDA_KEYS = {'l_a', 'l_g', 'l_i', 'l_sp', 'l_s', 'l_c'} + + +def _validate_weights( + gene_weight: float, + spot_weight: float, + spatial_weight: float, + ratios_weight: float, + sparsity_coef: float, + cluster_weight: float, +) -> None: + """ + Validate the optimisation loss weights, matching the original R/paper + formulas (see _default_weights_to_lambdas). + """ + raw = { + 'gene_weight': gene_weight, + 'spot_weight': spot_weight, + 'spatial_weight': spatial_weight, + 'ratios_weight': ratios_weight, + 'cluster_weight': cluster_weight + } + for name, value in raw.items(): + if value < 0: + raise ValueError(f"{name} must be >= 0, got {value}") + if sum(raw.values()) <= 0: + raise ValueError( + "At least one of gene_weight, spot_weight, spatial_weight, " + "ratios_weight must be > 0." + ) + if not (0 <= sparsity_coef <= 1): + raise ValueError(f"sparsity_coef must be in [0, 1], got {sparsity_coef}") + + +def _default_weights_to_lambdas( + weights: Dict[str, float], + *, + S: int, + G: int, + C: int, + max_size: int, + n_pairs: int, + has_pairs: bool, +) -> Dict[str, float]: + """ + The original R/paper formulas for turning loss weights into the actual + optimisation coefficients (lambdas). l_i already folds in + (1 - sparsity_coef) here, once, rather than re-applying it inside the + optimisation loop -- this is what makes a direct lambdas={'l_i': ...} + override in DOT.__init__() actually stick, instead of being silently + scaled down further by sparsity_coef at the point of use. + """ + sparsity_coef = weights['sparsity_coef'] + return { + 'l_a': weights['ratios_weight'] / max_size, + 'l_g': weights['gene_weight'] * S / G, + 'l_i': weights['spot_weight'] * (1 - sparsity_coef), + 'l_sp': weights['spot_weight'] * sparsity_coef / max_size, + 'l_c': weights['cluster_weight'] * S / C, + 'l_s': weights['spatial_weight'] * S / (max_size * n_pairs) if has_pairs else 0.0, + } + + +def _resolve_lambdas( + weights: Dict[str, float], + lambdas: Optional[Dict[str, float]], + weights_to_lambdas: Callable, + *, + S: int, + G: int, + C: int, + max_size: int, + n_pairs: int, + has_pairs: bool, +) -> Dict[str, float]: + """ + Resolve the final optimisation coefficients: weights_to_lambdas's + output, with any direct lambdas overrides applied on top. l_s is + forced to 0 (with a warning) if it would otherwise be nonzero despite + no spatial pairs existing, so opt_config['lambdas'] stays truthful + regardless of whether that came from the weight-derived default or an + explicit override. + """ + lambdas = lambdas or {} + unknown = set(lambdas) - _LAMBDA_KEYS + if unknown: + raise ValueError( + f"lambdas contains unknown keys {sorted(unknown)}, " + f"must be a subset of {sorted(_LAMBDA_KEYS)}" + ) + + computed_lambdas = weights_to_lambdas( + weights, S=S, G=G, C=C, max_size=max_size, n_pairs=n_pairs, has_pairs=has_pairs, + ) + missing = _LAMBDA_KEYS - set(computed_lambdas) + if missing: + raise ValueError( + f"weights_to_lambdas must return all of {sorted(_LAMBDA_KEYS)}, " + f"missing {sorted(missing)}" + ) + + final_lambdas = {**computed_lambdas, **lambdas} + if not has_pairs and final_lambdas['l_s'] != 0: + warnings.warn( + "l_s (spatial coherence) is nonzero, but no spatial neighbour " + "pairs were found in spatial -- it will have no effect on " + "fitting.", + UserWarning, + stacklevel=find_stack_level(), + ) + final_lambdas['l_s'] = 0.0 + + return final_lambdas + class DOT: """Deconvolution by Optimal Transport – GPU-optimised batched solver.""" def __init__( self, - spatial: Dict, - ref: Dict, + spatial: Union[Dict, AnnData], + ref: Union[Dict, AnnData], + mode: str = 'highres', + gene_weight: Optional[float] = None, + spot_weight: Optional[float] = None, + spatial_weight: Optional[float] = None, + ratios_weight: Optional[float] = None, + sparsity_coef: Optional[float] = None, + cluster_weight: Optional[float] = None, + max_size: Optional[int] = None, + min_size: int = 1, + lambdas: Optional[Dict[str, float]] = None, + weights_to_lambdas: Callable = _default_weights_to_lambdas, ls_solution: bool = True, batch_size: int = 500, device: Optional[str] = None, + verbose: bool = False, ): + """ + Parameters + ---------- + spatial, ref : Dict or AnnData + Output AnnData of setup_spatial() / setup_reference(). Dict is the + legacy shape and still supported, but deprecated. + mode : ``'highres'`` or ``'lowres'`` + Convenience preset supplying defaults for gene_weight, + spatial_weight, ratios_weight, sparsity_coef, and max_size. Any + of those passed explicitly overrides the preset's value for it. + gene_weight, spatial_weight, ratios_weight : float, optional + Relative weight of, respectively: gene-wise cosine fit, spatial + coherence, and reference abundance matching, matching the original + weight-to-lambda formulas (see weights_to_lambdas). + spot_weight : float, optional + Relative weight of spot-wise cosine fit. Defaults to 1.0 if the + resolved max_size == 1, else 0.25 (matches the original R/paper + behaviour) -- this default depends on max_size, not mode alone. + sparsity_coef : float, optional + In [0, 1]. Mixing ratio, within spot_weight's own share, between + spot-wise cosine fit (1 - sparsity_coef) and sparsity + (sparsity_coef). + cluster_weight : float, optional + Relative weight of cluster-wise cosine fit (does each reference + sub-cluster's spatially-implied expression profile still + resemble its own reference profile). Defaults to 0 (off), + matching the original R default. Not yet consumed by + fit() -- computed and shown in print_config(), but has no + effect on the optimisation until a later step lands. + max_size, min_size : int + Frank-Wolfe box constraint: min/max "cells" a spot's mass can + represent. + lambdas : dict, optional + Direct overrides for the actual optimisation coefficients + (any subset of 'l_a', 'l_g', 'l_i', 'l_sp', 'l_s', 'l_c'), applied + after weights_to_lambdas, bypassing weight-based scaling + entirely for just those terms. Everything not present here + still comes from the weights above. + weights_to_lambdas : callable, optional + Computes the coefficients from the resolved weights. Defaults + to the original R/paper formulas. Signature: + (weights: dict, *, S, G, C, max_size, n_pairs, has_pairs) -> dict, + returning all of 'l_a', 'l_g', 'l_i', 'l_sp', 'l_s', 'l_c'. + ls_solution : bool + Use ridge-regularised least-squares initialisation (recommended). + batch_size : int + Batch size for GPU processing. + device : str, optional + 'cuda' or 'cpu'. + verbose : bool + Print the resolved optimisation config (see print_config()). + """ + ref = _ref_to_anndata(ref) + spatial = _spatial_to_anndata(spatial) + # --- Gene alignment --- - spatial_genes = np.asarray(spatial['genes']) - ref_genes = np.asarray(ref['genes']) + spatial_genes = np.asarray(spatial.var_names) + ref_genes = np.asarray(ref.var_names) common_genes = np.intersect1d(spatial_genes, ref_genes) if len(common_genes) == 0: raise ValueError("No common genes found between spatial and reference data") + else: + print(f"{len(common_genes)} genes overlap between reference and spatial data") # Index both matrices in the same explicit order so columns correspond. sp_lookup = {gene: i for i, gene in enumerate(spatial_genes)} @@ -33,26 +300,56 @@ def __init__( sp_idx = np.fromiter((sp_lookup[g] for g in common_genes), dtype=np.int64) rf_idx = np.fromiter((rf_lookup[g] for g in common_genes), dtype=np.int64) - X_sp = spatial['X_sparse'][:, sp_idx] if issparse(spatial['X_sparse']) \ - else spatial['X_sparse'][:, sp_idx] - X_rf = ref['X_sparse'][:, rf_idx] if issparse(ref['X_sparse']) \ - else ref['X_sparse'][:, rf_idx] + self.spatial = spatial[:, sp_idx].copy() + self.ref = ref[:, rf_idx].copy() + self.device = device or 'cpu' + + # --- Optimisation config --- + if mode not in MODE_PRESETS: + raise ValueError(f"mode must be one of {list(MODE_PRESETS)}, got {mode!r}") + preset = MODE_PRESETS[mode] + + max_size = preset['max_size'] if max_size is None else max_size + gene_weight = preset['gene_weight'] if gene_weight is None else gene_weight + spot_weight = (1.0 if max_size == 1 else 0.25) if spot_weight is None else spot_weight + spatial_weight = preset['spatial_weight'] if spatial_weight is None else spatial_weight + ratios_weight = preset['ratios_weight'] if ratios_weight is None else ratios_weight + sparsity_coef = preset['sparsity_coef'] if sparsity_coef is None else sparsity_coef + cluster_weight = preset['cluster_weight'] if cluster_weight is None else cluster_weight + + _validate_weights( + gene_weight, spot_weight, spatial_weight, ratios_weight, sparsity_coef, cluster_weight, + ) - self.spatial = { - 'X_sparse': X_sp, - 'coords': spatial['coords'], - 'genes': common_genes, - 'device': device or 'cpu', + weights = { + 'gene_weight': gene_weight, + 'spot_weight': spot_weight, + 'spatial_weight': spatial_weight, + 'ratios_weight': ratios_weight, + 'sparsity_coef': sparsity_coef, + 'cluster_weight': cluster_weight, } - if 'pairs' in spatial: - self.spatial['pairs'] = spatial['pairs'] - - self.ref = { - 'X_sparse': X_rf, - 'clusters': ref['clusters'], - 'ratios': ref['ratios'], - 'genes': common_genes, - 'device': device or 'cpu', + + self._weights_to_lambdas = weights_to_lambdas + + S = self.spatial.n_obs + G = self.spatial.n_vars + C = self.ref.n_obs + has_pairs = 'dot_spatial_pairs' in self.spatial.obsp + n_pairs = triu(self.spatial.obsp['dot_spatial_pairs'], k=1).nnz if has_pairs else 0 + + self.opt_config = { + 'mode': mode, + 'weights': weights, + 'lambdas': _resolve_lambdas( + weights, lambdas, weights_to_lambdas, + S=S, G=G, C=C, max_size=max_size, n_pairs=n_pairs, has_pairs=has_pairs, + ), + 'lambda_overrides': set(lambdas) if lambdas else set(), + 'max_size': max_size, + 'min_size': min_size, + 'has_pairs': has_pairs, + 'n_pairs': n_pairs, } self.batch_size = batch_size @@ -63,13 +360,78 @@ def __init__( if ls_solution: self.solution = self._ls_solution() + if verbose: + self.print_config() + + # ------------------------------------------------------------------ + # Config inspection + # ------------------------------------------------------------------ + def print_config(self) -> None: + """Print the resolved optimisation configuration.""" + cfg = self.opt_config + w = cfg['weights'] + final_lambdas = cfg['lambdas'] + overrides = cfg['lambda_overrides'] + has_pairs = cfg['has_pairs'] + is_default = self._weights_to_lambdas is _default_weights_to_lambdas + + print(f"DOT optimisation configuration (mode={cfg['mode']!r}):") + print() + print("weights:") + for name in ('gene_weight', 'spot_weight', 'sparsity_coef', 'cluster_weight', 'spatial_weight', 'ratios_weight'): + print(f" {name:<15}: {w[name]:.4f}") + print(f" max_size/min_size: {cfg['max_size']}/{cfg['min_size']}") + if not is_default: + print( + f" (weights_to_lambdas: custom function {self._weights_to_lambdas.__name__!r} -- " + "weight/lambda mapping below is not the default formula)" + ) + print() + + # Recompute the weight/transform-derived (pre-override) lambdas + # fresh, purely for this display -- cheap (five floats), and avoids + # threading a second return value through _resolve_lambdas' already + # tested contract. + S = self.spatial.n_obs + G = self.spatial.n_vars + C = self.ref.n_obs + weight_derived = self._weights_to_lambdas( + w, S=S, G=G, C=C, max_size=cfg['max_size'], n_pairs=cfg['n_pairs'], has_pairs=has_pairs, + ) + + formulas = { + 'l_a': 'ratios_weight', + 'l_g': 'gene_weight', + 'l_i': 'spot_weight * (1 - sparsity_coef)', + 'l_sp': 'spot_weight * sparsity_coef', + 'l_c': 'cluster_weight', + 'l_s': 'spatial_weight', + } + + rows = [] + for key in ('l_g', 'l_i', 'l_sp', 'l_c', 'l_s', 'l_a'): + notes = [] + if key in overrides: + notes.append(f"user-set directly (weights alone: {weight_derived[key]:.4f})") + if key == 'l_s' and not has_pairs: + notes.append("inactive (no spatial pairs)") + row = {'term': key, 'lambda': final_lambdas[key], 'notes': ", ".join(notes)} + if is_default: + row['weight param'] = formulas[key] + rows.append(row) + + columns = ['term', 'weight param', 'lambda', 'notes'] if is_default \ + else ['term', 'lambda', 'notes'] + df = pd.DataFrame(rows, columns=columns).set_index('term') + print(df.to_string(float_format=lambda x: f"{x:.4f}")) + # ------------------------------------------------------------------ # Least-squares initialisation # ------------------------------------------------------------------ def _ls_solution(self, lambda_ridge: float = 100.0) -> np.ndarray: """Ridge-regularised LS init. Exploits sparsity when possible.""" - X_ref = self.ref['X_sparse'] - X_sp = self.spatial['X_sparse'] + X_ref = self.ref.X + X_sp = self.spatial.X if issparse(X_ref): X_ref_d = X_ref.toarray().astype(np.float32) @@ -99,9 +461,6 @@ def _ls_solution(self, lambda_ridge: float = 100.0) -> np.ndarray: # ------------------------------------------------------------------ def fit( self, - mode: str = 'highres', - ratios_weight: float = 0.0, - max_spot_size: int = 20, iterations: int = 100, gap_threshold: float = 0.01, verbose: bool = False, @@ -113,13 +472,11 @@ def fit( """ Run DOT optimisation. + The optimisation objective itself (mode / weights / sparsity_coef / + max_size / lambdas) is configured at DOT() construction time. + Parameters ---------- - mode : ``'highres'`` or ``'lowres'`` - ratios_weight : float - Weight for matching reference cell-type abundances. - max_spot_size : int - Max cells per spot (lowres mode). iterations : int Frank-Wolfe iterations. gap_threshold : float @@ -131,22 +488,11 @@ def fit( use_mixed_precision : bool Use float16 intermediates on GPU (saves memory). """ - if mode == 'highres': - sparsity_coef, max_size = 0.6, 1 - elif mode == 'lowres': - sparsity_coef, max_size = 0.4, max_spot_size - else: - raise ValueError("mode must be 'highres' or 'lowres'") - start_iter = 1 if resume_from is not None: start_iter = self._load_checkpoint(resume_from, verbose) self._run_optimisation( - ratios_weight=ratios_weight, - sparsity_coef=sparsity_coef, - max_size=max_size, - min_size=1, iterations=iterations, gap_threshold=gap_threshold, verbose=verbose, @@ -163,12 +509,11 @@ def fit( @torch.no_grad() def _run_optimisation( self, - ratios_weight, sparsity_coef, max_size, min_size, iterations, gap_threshold, verbose, checkpoint_dir, checkpoint_freq, start_iteration, use_mixed_precision=False, ): - device_str = self.ref['device'] + device_str = self.device use_gpu = device_str == 'cuda' and torch.cuda.is_available() device = torch.device(device_str if use_gpu else 'cpu') @@ -180,48 +525,63 @@ def _run_optimisation( # ============================================================ # 1. Prepare data on CPU # ============================================================ - X_ref_np = self.ref['X_sparse'].toarray().astype(np.float32) \ - if issparse(self.ref['X_sparse']) else np.asarray(self.ref['X_sparse'], dtype=np.float32) - X_sp_np = self.spatial['X_sparse'].toarray().astype(np.float32) \ - if issparse(self.spatial['X_sparse']) else np.asarray(self.spatial['X_sparse'], dtype=np.float32) + X_ref_np = self.ref.X.toarray().astype(np.float32) \ + if issparse(self.ref.X) else np.asarray(self.ref.X, dtype=np.float32) + X_sp_np = self.spatial.X.toarray().astype(np.float32) \ + if issparse(self.spatial.X) else np.asarray(self.spatial.X, dtype=np.float32) S, G = X_sp_np.shape C = X_ref_np.shape[0] - cell_types = list(self.ref['clusters'].keys()) + clusters = _clusters_at_level(self.ref) + cell_types = list(clusters.keys()) K = len(cell_types) # Cluster → major type mapping (vectorised) cluster_to_major = np.zeros(C, dtype=np.int64) cluster_indices_list = [] # list of np arrays per major type for k, ct in enumerate(cell_types): - idx = np.asarray(self.ref['clusters'][ct]) + idx = np.asarray(clusters[ct]) cluster_to_major[idx] = k cluster_indices_list.append(idx) - sc_ratios = np.array([self.ref['ratios'][ct] for ct in cell_types], dtype=np.float32) + sc_ratios = np.array([self.ref.uns['ratios'][ct] for ct in cell_types], dtype=np.float32) sc_ratios /= sc_ratios.sum() + max_size = self.opt_config['max_size'] + min_size = self.opt_config['min_size'] + # Only needed below for warm-start rescaling -- the loop itself + # reads l_i/l_sp, which already fold sparsity_coef in (see + # _default_weights_to_lambdas). + sparsity_coef = self.opt_config['weights']['sparsity_coef'] + r_st = np.full(S, 0.9 * min_size + 0.1 * max_size, dtype=np.float32) n_st = r_st.sum() r_sc = sc_ratios * n_st r_sc_ex = r_sc[cluster_to_major] - # Loss weights - inner = [1.0, 0.25 if max_size > 1 else 1.0, 0.0, 0.01] - l_a = ratios_weight / max_size - l_g = inner[0] * S / G - l_i = inner[1] - l_sp = l_i * sparsity_coef / max_size - - has_pairs = 'pairs' in self.spatial and self.spatial['pairs'] is not None - if has_pairs: - l_s = inner[3] * S / (max_size * len(self.spatial['pairs']['i'])) - pairs_i_np = self.spatial['pairs']['i'].astype(np.int64) - pairs_j_np = self.spatial['pairs']['j'].astype(np.int64) - pairs_w_np = self.spatial['pairs']['w'].astype(np.float32) - else: - l_s = 0.0 + # Loss coefficients, resolved once at DOT() construction time. + lambdas = self.opt_config['lambdas'] + l_a = lambdas['l_a'] + l_g = lambdas['l_g'] + l_i = lambdas['l_i'] + l_sp = lambdas['l_sp'] + l_c = lambdas['l_c'] + l_s = lambdas['l_s'] + + use_pairs = self.opt_config['has_pairs'] and l_s > 0 + if use_pairs: + pairs_mat = triu(self.spatial.obsp['dot_spatial_pairs'], k=1).tocoo() + pairs_i_np = pairs_mat.row.astype(np.int64) + pairs_j_np = pairs_mat.col.astype(np.int64) + pairs_w_np = pairs_mat.data.astype(np.float32) + + if l_c > 0: + w_sc = np.zeros(C, dtype=np.float32) + for k in range(K): + idx = cluster_indices_list[k] + w_sc[idx] = r_sc[k] / len(idx) + w_sc = w_sc / w_sc.sum() * C # ============================================================ # 2. Move to device ONCE @@ -231,12 +591,16 @@ def _run_optimisation( X_ref_norm = F.normalize(X_ref, p=2, dim=1) # C × G (L2-normed rows) X_sp_row_norm = F.normalize(X_sp, p=2, dim=1) # S × G (spot-wise cosine + linear sparsity) X_sp_col_norm = F.normalize(X_sp, p=2, dim=0) # S × G (gene-wise cosine) - del X_sp # free – raw spatial no longer needed + if l_c > 0: + X_sp_norms = X_sp.norm(dim=1, keepdim=True) # S × 1 (raw magnitude, cluster-wise cosine) + del X_sp # free – reconstructed per-batch below if needed c2m = torch.from_numpy(cluster_to_major).to(device) r_sc_t = torch.from_numpy(r_sc).to(device) sc_ratios_t = torch.from_numpy(sc_ratios).to(device) r_st_t = torch.from_numpy(r_st).to(device) + if l_c > 0: + w_sc_t = torch.from_numpy(w_sc).to(device) # Pre-build scatter indices for cluster→major aggregation # cluster_scatter[c] = k (same as c2m, but we keep both for clarity) @@ -249,7 +613,7 @@ def _run_optimisation( torch.from_numpy(cluster_indices_list[k]).to(device) ) - if has_pairs: + if use_pairs: p_i = torch.from_numpy(pairs_i_np).to(device) p_j = torch.from_numpy(pairs_j_np).to(device) p_w = torch.from_numpy(pairs_w_np).to(device) * 0.5 / np.log(2) @@ -349,7 +713,7 @@ def _run_optimisation( dcosine_st = 0.0 dcosine_lin = 0.0 - need_spot = (sparsity_coef < 1 and l_i > 0) or l_sp > 0 + need_spot = l_i > 0 or l_sp > 0 if need_spot: for b in range(n_batches): s0 = b * batch @@ -365,7 +729,7 @@ def _run_optimisation( st_xt = Yt_b.T @ X_ref # -- spot-wise cosine -- - if sparsity_coef < 1 and l_i > 0: + if l_i > 0: norms = st_xt.norm(dim=1, keepdim=True).clamp(min=1e-10) st_xt_n = st_xt / norms @@ -375,8 +739,7 @@ def _run_optimisation( di_sqrt = _sqrt_env(di) dcosine_st += di_sqrt.sum().item() - coef = l_i * (1 - sparsity_coef) - st_de = coef * (Xsp_b_n - st_xt_n * csi.unsqueeze(1)) \ + st_de = l_i * (Xsp_b_n - st_xt_n * csi.unsqueeze(1)) \ * d_i_grad.unsqueeze(1) / norms if compute_dtype == torch.float16: @@ -394,6 +757,51 @@ def _run_optimisation( Dt[:, s0:s1].add_(lin_d, alpha=l_sp) dcosine_lin += (Yt_b * lin_d).sum().item() + # ============ CLUSTER-WISE COSINE (batched) ============ + dcosine_sc = 0.0 + if l_c > 0: + # pass 1: accumulate each cluster's spatially-implied + # profile in chunks -- raw X_sp isn't kept resident, so + # each batch's values are reconstructed from the + # already-normalized rows and their stored magnitude + sc_xt = torch.zeros(C, G, device=device, dtype=torch.float32) # C × G + for b in range(n_batches): + s0 = b * batch + s1 = min(s0 + batch, S) + Yt_b = Yt[:, s0:s1] # C × b + Xsp_b = X_sp_row_norm[s0:s1] * X_sp_norms[s0:s1] # b × G (reconstructed raw) + + if compute_dtype == torch.float16: + sc_xt += (Yt_b.half() @ Xsp_b.half()).float() + else: + sc_xt += Yt_b @ Xsp_b + + sc_xt_inorms = sc_xt.norm(dim=1, keepdim=True).clamp(min=1e-10) + sc_xt_n = sc_xt / sc_xt_inorms + + csc = (sc_xt_n * X_ref_norm).sum(dim=1) + dc = (1 - csc).clamp(min=0) + dc_coefs = w_sc_t / sc_xt_inorms.squeeze(1) + + dc_coefs = dc_coefs * _sqrt_env_grad(dc) + dc_sqrt = _sqrt_env(dc) + dcosine_sc = (dc_sqrt * w_sc_t).sum().item() + + sc_de = l_c * (X_ref_norm - sc_xt_n * csc.unsqueeze(1)) * dc_coefs.unsqueeze(1) + + # pass 2: broadcast the small per-cluster correction back + # onto every spot, again reconstructing raw values batch + # by batch + for b in range(n_batches): + s0 = b * batch + s1 = min(s0 + batch, S) + Xsp_b = X_sp_row_norm[s0:s1] * X_sp_norms[s0:s1] # b × G + + if compute_dtype == torch.float16: + Dt[:, s0:s1] -= (Xsp_b.half() @ sc_de.half().T).float().T + else: + Dt[:, s0:s1] -= (Xsp_b @ sc_de.T).T + # ============ GENE-WISE COSINE (chunked) ============ dcosine_g = 0.0 if l_g > 0: @@ -419,7 +827,7 @@ def _run_optimisation( # ============ SPATIAL COHERENCE (vectorised) ============ d_s = 0.0 - if l_s > 0 and has_pairs: + if use_pairs: # Ytk is K × S # Vectorised over all pairs at once n_pairs = p_i.shape[0] @@ -470,8 +878,9 @@ def _run_optimisation( Yt_h[kk, torch.arange(S, device=device)] = fill_val # ---- Objective + gap ---- - ft = (l_i * (1 - sparsity_coef) * dcosine_st + ft = (l_i * dcosine_st + l_sp * dcosine_lin + + l_c * dcosine_sc + l_g * dcosine_g + l_s * d_s + l_a * ratio_err) @@ -586,18 +995,46 @@ def _load_checkpoint(self, path, verbose): # ------------------------------------------------------------------ # Results # ------------------------------------------------------------------ - def get_weights(self, normalize: bool = True) -> np.ndarray: + def get_weights(self, level: Optional[str] = None, normalize: bool = True) -> pd.DataFrame: + """ + Per-spot cell-type weights, aggregated to the given cell type level. + + Parameters + ---------- + level : str, optional + A column name from ref.uns['level_keys']. Defaults to the + coarsest level. Finer levels are computed lazily on demand + from the finest-grained solution. + normalize : bool + Normalize each spot's weights to sum to 1. + + Returns + ------- + pd.DataFrame + Spots (index) x cell types at the requested level (columns). + """ if self.weights is None: raise ValueError("Not fitted yet – call fit() first.") - w = self.weights.copy() + clusters = _clusters_at_level(self.ref, level) + cell_types = list(clusters.keys()) + + if level is None or level == self.ref.uns['level_keys'][0]: + w = self.weights.copy() + else: + w = np.zeros((self.solution.shape[1], len(cell_types)), dtype=np.float32) + for k, ct in enumerate(cell_types): + w[:, k] = self.solution[clusters[ct]].sum(axis=0) + if normalize: rs = w.sum(axis=1, keepdims=True) rs[rs == 0] = 1 - w /= rs - return w + w = w / rs + + return pd.DataFrame(w, index=self.spatial.obs_names, columns=cell_types) - def get_cell_types(self) -> list: - return list(self.ref['clusters'].keys()) + def get_cell_types(self, level: Optional[str] = None) -> list: + """Cell-type labels at the given hierarchy level (defaults to coarsest).""" + return list(_clusters_at_level(self.ref, level).keys()) # ====================================================================== @@ -605,10 +1042,9 @@ def get_cell_types(self) -> list: # ====================================================================== def _safe_log2(x: torch.Tensor) -> torch.Tensor: - """log2 that maps 0 → 0 and clips -inf.""" - out = torch.log2(x.clamp(min=1e-10)) - out = torch.nan_to_num(out, nan=0.0, posinf=0.0, neginf=-20.0) - return out + """log2 with NaN/±inf replaced, matching R's safelog2: NaN -> 0, ±inf -> -20.""" + out = torch.log2(x) + return torch.nan_to_num(out, nan=0.0, posinf=-20.0, neginf=-20.0) _ENV_MIN = 1e-2 diff --git a/dotpy/ds.py b/dotpy/ds.py new file mode 100644 index 0000000..26a46d4 --- /dev/null +++ b/dotpy/ds.py @@ -0,0 +1,127 @@ +"""Synthetic datasets for testing and examples.""" + +from typing import Tuple + +import numpy as np + + +def toy_reference( + n_major_types: int = 2, + n_subtypes_per_major: int = 3, + n_cells_per_subtype: int = 50, + n_genes: int = 100, + n_major_markers: int = 10, + n_subtype_markers: int = 5, + major_effect: float = 20.0, + subtype_effect: float = 8.0, + background_range: Tuple[int, int] = (1, 10), + major_prefix: str = "Major", + subtype_prefix: str = "Sub", + name_sep: str = "_", + random_state: int = 0, +) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: + """ + Generate synthetic expression data with a known, two-level (major type + to subtype) marker structure, for testing subtype-discovery code + against ground truth. + + Parameters + ---------- + n_major_types : int + Number of major types to generate. + n_subtypes_per_major : int + Number of subtypes within each major type. + n_cells_per_subtype : int + Number of cells per subtype (uniform across subtypes). + n_genes : int + Total number of genes. Must be large enough to fit every major-type + and subtype marker block without overlap. + n_major_markers : int + Number of marker genes per major type. + n_subtype_markers : int + Number of marker genes per subtype. + major_effect : float + Amount added to a major type's marker genes, for its own cells. + subtype_effect : float + Amount added to a subtype's marker genes, for its own cells. + background_range : tuple of int + (low, high) bounds of the background noise, before any marker + effect is added. low is inclusive, high is exclusive, matching + numpy's Generator.integers. low must be < high. + major_prefix : str + Prefix for major type names, e.g. "Major" -> "Major0". + subtype_prefix : str + Prefix for the subtype part of a subtype name, e.g. "Sub" -> + "Major0_Sub1". + name_sep : str + Separator between a major type name and its subtype prefix. + name_sep and subtype_prefix cannot both be empty. + random_state : int + Random seed for reproducibility. + + Returns + ------- + X : np.ndarray, shape (n_cells, n_genes) + Synthetic expression matrix. Each major type has its own block of + marker genes elevated for all its cells. Within each major type, + each subtype has its own, smaller block of marker genes elevated + only for its own cells. Everywhere else is uniform background + noise. + major_types : np.ndarray, shape (n_cells,) + Major type label for each cell. + subtypes : np.ndarray, shape (n_cells,) + Ground-truth subtype label for each cell. Names are unique across + major types, e.g. "Major0_Sub1". This makes the output usable both + as ground truth for validating discovered subtypes and as a + predefined subtype annotation in its own right. + """ + n_major_marker_genes = n_major_types * n_major_markers + n_subtype_marker_genes = n_major_types * n_subtypes_per_major * n_subtype_markers + if n_major_marker_genes + n_subtype_marker_genes > n_genes: + raise ValueError( + f"n_genes={n_genes} too small, need " + f">= {n_major_marker_genes + n_subtype_marker_genes}." + ) + + for param_name, value in [ + ("major_prefix", major_prefix), ("subtype_prefix", subtype_prefix), ("name_sep", name_sep) + ]: + if not isinstance(value, str): + raise TypeError(f"{param_name} must be a str, got {type(value).__name__}.") + + if name_sep == "" and subtype_prefix == "": + raise ValueError("name_sep and subtype_prefix cannot both be empty.") + + if len(background_range) != 2: + raise ValueError(f"background_range must have 2 values, got {len(background_range)}.") + background_low, background_high = background_range + if background_low >= background_high: + raise ValueError(f"background_range low ({background_low}) must be < high ({background_high}).") + + rng = np.random.default_rng(random_state) + + n_cells = n_major_types * n_subtypes_per_major * n_cells_per_subtype + X = rng.integers(background_low, background_high, size=(n_cells, n_genes)).astype(np.float64) + major_types = np.empty(n_cells, dtype=object) + subtypes = np.empty(n_cells, dtype=object) + + next_marker_gene = 0 + row = 0 + for m in range(n_major_types): + major_name = f"{major_prefix}{m}" + major_genes = slice(next_marker_gene, next_marker_gene + n_major_markers) + next_marker_gene += n_major_markers + + for s in range(n_subtypes_per_major): + subtype_name = f"{major_name}{name_sep}{subtype_prefix}{s}" + subtype_genes = slice(next_marker_gene, next_marker_gene + n_subtype_markers) + next_marker_gene += n_subtype_markers + + block = slice(row, row + n_cells_per_subtype) + X[block, major_genes] += major_effect + X[block, subtype_genes] += subtype_effect + major_types[block] = major_name + subtypes[block] = subtype_name + row += n_cells_per_subtype + + return X, major_types, subtypes diff --git a/dotpy/preprocessing.py b/dotpy/preprocessing.py index 5923a84..f0d0de0 100644 --- a/dotpy/preprocessing.py +++ b/dotpy/preprocessing.py @@ -4,14 +4,101 @@ Memory-efficient implementation with sparse matrix support. """ +import heapq +import warnings + import numpy as np -from typing import Optional, Dict +import pandas as pd +from typing import Callable, List, Optional, Dict, Tuple, Union from anndata import AnnData from scipy.sparse import issparse, csr_matrix, vstack from sklearn.cluster import KMeans, MiniBatchKMeans from sklearn.preprocessing import normalize as sk_normalize from sklearn.neighbors import NearestNeighbors -import scanpy as sc + +from ._exceptions import find_stack_level + +from ._exceptions import find_stack_level + + +def check_subtype_consistency( + obs: pd.DataFrame, + level_keys: List[str], + depth: Optional[int] = None, +) -> None: + """ + Validate that a set of hierarchical cell type annotation columns forms a unique + nested hierarchy: each label at a finer level maps to exactly one + label at coarser levels. + + Parameters + ---------- + obs : pd.DataFrame + Table containing the annotation columns to check, e.g. adata.obs. + level_keys : list of str + Column names in obs, ordered from coarsest (e.g. major cell type) + to finest (e.g. cell state). Must have at least 2 entries. + depth : int, optional + Number of levels (counted from the coarsest) to validate; must be + >= 2, or None to validate every level given (default). + + Raises + ------ + ValueError + If fewer than 2 levels are given/requested, if a level_keys entry + is missing from obs, if a checked column contains missing (NaN) + values, or if a label at one level maps to more than one label at + the immediately coarser level. + """ + if len(level_keys) < 2: + raise ValueError( + f"level_keys must have at least 2 entries to check hierarchy " + f"consistency, got {len(level_keys)}." + ) + + if depth is None: + depth = len(level_keys) + if not (2 <= depth <= len(level_keys)): + raise ValueError( + f"depth must be between 2 and len(level_keys)={len(level_keys)}, " + f"got {depth}." + ) + + checked_keys = level_keys[:depth] + + missing = [key for key in checked_keys if key not in obs.columns] + if missing: + raise ValueError( + f"level_keys entries not found in obs.columns: {missing}. " + f"Available columns: {list(obs.columns)}" + ) + + for key in checked_keys: + if obs[key].isna().any(): + n_na = int(obs[key].isna().sum()) + raise ValueError( + f"Column '{key}' has {n_na} missing (NaN) value(s). " + f"check_subtype_consistency requires complete annotations " + f"for every level being checked." + ) + + for coarse_key, fine_key in zip(checked_keys[:-1], checked_keys[1:]): + n_parents = obs.groupby(fine_key, observed=True)[coarse_key].nunique() + offenders = n_parents[n_parents > 1] + if len(offenders) > 0: + examples = [] + for label in offenders.index[:5]: + parents = sorted( + obs.loc[obs[fine_key] == label, coarse_key].unique().tolist() + ) + examples.append(f" '{label}' -> {parents}") + more = f" (and {len(offenders) - 5} more)" if len(offenders) > 5 else "" + raise ValueError( + f"Inconsistent hierarchy between '{coarse_key}' (coarser) and " + f"'{fine_key}' (finer): {len(offenders)} label(s) in '{fine_key}' " + f"map to more than one '{coarse_key}' label{more}:\n" + + "\n".join(examples) + ) def _select_kmeans_genes( @@ -54,7 +141,7 @@ def _kmeans_subcluster( X_ct: np.ndarray, gene_indices: np.ndarray, K: int, - min_frac: float = 0.025, + min_cells: float, random_state: int = 42 ) -> np.ndarray: """K-means clustering with small cluster filtering.""" @@ -81,9 +168,8 @@ def _kmeans_subcluster( labels = km.fit_predict(X_subset) - n_cells = len(labels) unique_labels, counts = np.unique(labels, return_counts=True) - noise_mask = counts < (min_frac * n_cells) + noise_mask = counts < min_cells noise_clusters = unique_labels[noise_mask] if len(noise_clusters) > 0: @@ -94,254 +180,1102 @@ def _kmeans_subcluster( return labels -def _get_de_genes_r_style( - centroids: np.ndarray, - max_genes: int, - verbose: bool = False -) -> np.ndarray: - """Select DE genes using R's median rank scoring method.""" - if centroids.shape[1] <= max_genes: - return np.arange(centroids.shape[1]) +def _default_k_heuristic(n_cells: int) -> float: + """R-ported heuristic: a major type's natural subtype count from its cell count.""" + return 2 * np.log(n_cells) - 7 - C, G = centroids.shape - if issparse(centroids): - centroids_dense = centroids.toarray() + 1e-9 - else: - centroids_dense = centroids.copy() + 1e-9 +def _expected_k(n_cells: int, subcluster_size: int, k_heuristic: Callable[[int], float]) -> int: + """Heuristic subtype count for n_cells, capped at subcluster_size and floored at 1.""" + return min(subcluster_size, max(1, int(np.round(k_heuristic(n_cells))))) - gene_scores = np.zeros((C, G), dtype=np.float32) - if verbose: - print(f"Computing gene scores for {C} clusters, {G} genes...") +def _huntington_hill_apportion(weights: Dict[str, int], budget: int) -> Dict[str, int]: + """ + Allocate `budget` discrete units across groups in `weights`, proportional + to each group's weight, guaranteeing every group at least 1 unit. - for i in range(C): - this_ct = np.tile(centroids_dense[i:i+1, :], (C-1, 1)) - other_ct = np.delete(centroids_dense, i, axis=0) - logfc = np.log(this_ct / other_ct) + Implements the Huntington-Hill method + (https://en.wikipedia.org/wiki/Huntington%E2%80%93Hill_method). + """ + if not weights: + raise ValueError("weights must not be empty.") + if any(w <= 0 for w in weights.values()): + raise ValueError(f"All weights must be positive, got {weights}.") + if budget < len(weights): + raise ValueError( + f"budget ({budget}) must be >= the number of groups " + f"({len(weights)}) so every group can get its guaranteed unit." + ) - ranks = np.empty_like(logfc) - for j in range(C-1): - ranks[j, :] = G - np.argsort(np.argsort(logfc[j, :])) + seats = {label: 1 for label in weights} - gene_scores[i, :] = np.median(ranks, axis=0) + def priority(label): + s = seats[label] + return weights[label] / (s * (s + 1)) ** 0.5 - min_scores = gene_scores.min(axis=0) - top_genes = np.argsort(min_scores)[:max_genes] + heap = [(-priority(label), label) for label in weights] + heapq.heapify(heap) - if verbose: - print(f"Selected {len(top_genes)} DE genes using R-style ranking") + for _ in range(budget - len(weights)): + _, label = heapq.heappop(heap) + seats[label] += 1 + heapq.heappush(heap, (-priority(label), label)) - return top_genes + return seats -def _aggregate_reference( - X: np.ndarray, - annotations: np.ndarray, - cluster_size: int, - th_inner_logfold: float = 0.75, - random_state: int = 42, - verbose: bool = False -) -> Dict: - """Aggregate reference data with R-compatible algorithm.""" - np.random.seed(random_state) - - major_types = np.unique(annotations) - n_genes = X.shape[1] - - if verbose: - print("Computing major centroids...") +def _capped_group_centroids(X, group_labels: np.ndarray, cap: int = 1000) -> Tuple[np.ndarray, Dict]: + """Per-group centroid (subsample capped at `cap` cells) and true abundance ratio.""" + groups = np.unique(group_labels) candidate_indices = [] candidate_types = [] - - for ct in major_types: - ct_idx = np.where(annotations == ct)[0] - - if len(ct_idx) > 1000: - ct_idx = np.random.choice(ct_idx, 1000, replace=False) - - candidate_indices.extend(ct_idx) - candidate_types.extend([ct] * len(ct_idx)) - + for g in groups: + g_idx = np.where(group_labels == g)[0] + if len(g_idx) > cap: + g_idx = np.random.choice(g_idx, cap, replace=False) + candidate_indices.extend(g_idx) + candidate_types.extend([g] * len(g_idx)) candidate_indices = np.array(candidate_indices) candidate_types = np.array(candidate_types) - major_centroids_list = [] - major_ratios = {} - - for ct in major_types: - ct_mask = candidate_types == ct - + centroids_list = [] + ratios = {} + for g in groups: + g_mask = candidate_types == g if issparse(X): - centroid = np.asarray(X[candidate_indices[ct_mask]].mean(axis=0)).flatten() + centroid = np.asarray(X[candidate_indices[g_mask]].mean(axis=0)).flatten() else: - centroid = X[candidate_indices[ct_mask]].mean(axis=0) + centroid = X[candidate_indices[g_mask]].mean(axis=0) + centroids_list.append(centroid) + ratios[g] = int((group_labels == g).sum()) + total = sum(ratios.values()) + ratios = {k: v / total for k, v in ratios.items()} - major_centroids_list.append(centroid) - major_ratios[ct] = int((annotations == ct).sum()) + return np.array(centroids_list), ratios - total = sum(major_ratios.values()) - major_ratios = {k: v / total for k, v in major_ratios.items()} - major_centroids = np.array(major_centroids_list) +def kmeans_define_subtypes( + X, + annotations: np.ndarray, + subcluster_size: int = 10, + th_inner_logfold: float = 0.75, + min_frac: float = 0.025, + name_sep: str = "_", + random_state: int = 42, + verbose: bool = False, + k_heuristic: Callable[[int], float] = _default_k_heuristic, + max_cells_per_type: int = 10000, + max_genes: int = 500, + annotations_key: str = "major", + subtype_key: str = "subtype", +) -> pd.DataFrame: + """ + Assign each cell a subtype label by sub-clustering each major type with + k-means. - if cluster_size <= 1: - if issparse(X): - sub_centroids = vstack([csr_matrix(c) for c in major_centroids_list]) - else: - sub_centroids = major_centroids + Parameters + ---------- + X : array-like or sparse matrix + Cell x gene expression matrix. + annotations : np.ndarray + Major type label for each cell (same length as X's first axis). + subcluster_size : int + Maximum number of subtypes per major type -- always applied as a + cap on top of k_heuristic's result, regardless of which heuristic + is used. Must be >= 1. + th_inner_logfold : float + Log-fold threshold for selecting genes used to find subtypes. + min_frac : float + Minimum fraction of a major type's cells a subtype must contain to + be kept; smaller ones are dropped as noise. 0 disables dropping. + name_sep : str + Separator used to build subtype names from their major type and an + index, e.g. f"{major_type}{name_sep}{i}". + random_state : int + Random seed for reproducibility. + verbose : bool + Print progress messages. + k_heuristic : callable + Maps a major type's (sampled) cell count to its natural subtype + count, before the subcluster_size cap is applied. Defaults to the + R-ported heuristic ``2 * log(n_cells) - 7``. + max_cells_per_type : int + Cap on how many of a major type's cells are used for clustering; + types larger than this are randomly subsampled. Cells left out are + not assigned a subtype. Must be >= 1. + max_genes : int + Gene narrowing only runs when the input has more than this many + genes, and keeps at most this many when it does. + annotations_key : str + Column name for the input annotations in the returned DataFrame. + subtype_key : str + Column name for the newly assigned subtypes in the returned + DataFrame. - clusters = {ct: [i] for i, ct in enumerate(major_types)} + Returns + ------- + pd.DataFrame + Two columns, annotations_key and subtype_key, one row per cell in + X/annotations. subtype_key is None for cells not assigned to any + subtype. + + Raises + ------ + ValueError + If subcluster_size or max_cells_per_type is < 1, or if + annotations_key equals subtype_key. + + Warns + ----- + UserWarning + If a major type has too few cells to sub-cluster, if all of a major + type's candidate subtypes were dropped as noise, or if any cells + were dropped as noise. + """ + if subcluster_size < 1: + raise ValueError(f"subcluster_size must be >= 1, got {subcluster_size}.") + if max_cells_per_type < 1: + raise ValueError(f"max_cells_per_type must be >= 1, got {max_cells_per_type}.") + if annotations_key == subtype_key: + raise ValueError(f"annotations_key and subtype_key must differ, both were '{annotations_key}'.") - return { - 'major_centroids': major_centroids, - 'major_ratios': major_ratios, - 'sub_centroids': sub_centroids, - 'clusters': clusters - } + np.random.seed(random_state) - if verbose: - print(f"Sub-clustering {len(major_types)} cell types...") + n_genes = X.shape[1] + labels_out = np.full(X.shape[0], None, dtype=object) - sub_centroids_list = [] - clusters = {} + major_types = np.unique(annotations) - for ct_idx, ct in enumerate(major_types): - ct_mask = annotations == ct - ct_indices = np.where(ct_mask)[0] + # Coarse per-major-type centroids, used only to pick marker-like genes + # for clustering below -- skipped entirely if there aren't enough genes + # to need narrowing down first. + major_centroids = None + major_ratios = None + if n_genes > max_genes: + major_centroids, major_ratios = _capped_group_centroids(X, annotations) + for ct_pos, ct in enumerate(major_types): + ct_indices = np.where(annotations == ct)[0] + + # Too few cells to form even one meaningful sub-cluster. if len(ct_indices) <= 1: + warnings.warn( + f"Major type '{ct}' has only {len(ct_indices)} cell(s); at " + f"least 2 are needed to sub-cluster, so it has no subtype " + f"assigned.", + UserWarning, + stacklevel=find_stack_level(), + ) continue - if len(ct_indices) > 10000: - ct_indices = np.random.choice(ct_indices, 10000, replace=False) + # Bound clustering cost; cells left out simply get no subtype. + sampled_indices = ct_indices + if len(ct_indices) > max_cells_per_type: + sampled_indices = np.random.choice(ct_indices, max_cells_per_type, replace=False) - X_ct = X[ct_indices] + X_ct = X[sampled_indices] - K = min(cluster_size, max(1, int(np.round(2 * np.log(len(ct_indices)) - 7)))) + # Heuristic target subtype count, capped at subcluster_size. + K = _expected_k(len(sampled_indices), subcluster_size, k_heuristic) + # Skip clustering for too few cells for more than one subtype if K <= 1: - if issparse(X_ct): - centroid = np.asarray(X_ct.mean(axis=0)).flatten() - else: - centroid = X_ct.mean(axis=0) - sub_centroids_list.append(centroid) - clusters[ct] = [len(sub_centroids_list) - 1] + labels_out[sampled_indices] = f"{ct}{name_sep}0" + if verbose: + print(f" {ct}: single subtype '{ct}{name_sep}0' ({len(sampled_indices)} cells)") continue + # Narrow down to marker-like genes before clustering kmeans_genes = np.arange(n_genes) - - if n_genes > 500: + if n_genes > max_genes: kmeans_genes = _select_kmeans_genes( - ct_centroid=major_centroids[ct_idx], + ct_centroid=major_centroids[ct_pos], major_centroids=major_centroids, major_ratios=major_ratios, ct_name=ct, th_logfold=th_inner_logfold, - max_genes=500 + max_genes=max_genes, ) if verbose: - print(f" Clustering {len(ct_indices)} {ct} cells into ~{K} clusters " - f"(using {len(kmeans_genes)} genes)...") + print(f" Clustering {len(sampled_indices)} '{ct}' cells into " + f"~{K} subtypes (using {len(kmeans_genes)} genes)...") labels = _kmeans_subcluster( X_ct=X_ct, gene_indices=kmeans_genes, K=K, - min_frac=0.025, - random_state=random_state + min_cells=min_frac * len(sampled_indices), + random_state=random_state, ) valid_mask = labels >= 0 + n_dropped_cells = int((~valid_mask).sum()) + # Fall back to one subtype for the whole (sampled) major type. + # if every candidate sub-cluster was too small to keep if valid_mask.sum() == 0: - if issparse(X_ct): - centroid = np.asarray(X_ct.mean(axis=0)).flatten() - else: - centroid = X_ct.mean(axis=0) - sub_centroids_list.append(centroid) - clusters[ct] = [len(sub_centroids_list) - 1] + labels_out[sampled_indices] = f"{ct}{name_sep}0" + warnings.warn( + f"All candidate subtypes of major type '{ct}' were below " + f"min_frac={min_frac} and dropped as noise; assigned a " + f"single subtype '{ct}{name_sep}0' instead.", + UserWarning, + stacklevel=find_stack_level(), + ) continue - X_ct_valid = X_ct[valid_mask] - labels_valid = labels[valid_mask] + # Name and assign each surviving sub-cluster; noise cells (-1) are + # skipped, leaving them None in labels_out. + unique_labels = np.unique(labels[valid_mask]) + for i, sc_label in enumerate(unique_labels): + cell_mask = labels == sc_label + labels_out[sampled_indices[cell_mask]] = f"{ct}{name_sep}{i}" + + if n_dropped_cells > 0: + warnings.warn( + f"Dropped {n_dropped_cells} cell(s) from major type '{ct}' " + f"as noise (subtype(s) below min_frac={min_frac} of its " + f"{len(sampled_indices)} sampled cells).", + UserWarning, + stacklevel=find_stack_level(), + ) + + return pd.DataFrame({ + annotations_key: annotations, + subtype_key: pd.Series(labels_out, dtype=object), + }) + + +def plan_subtype_refinement( + obs: pd.DataFrame, + level_keys: List[str], + subcluster_size: int = 10, + min_frac: float = 0.025, + k_heuristic: Callable[[int], float] = _default_k_heuristic, + return_counts: bool = False, + verbose: bool = False, +) -> Optional[pd.DataFrame]: + """ + Diagnose over-/under-clustering in a predefined subtype hierarchy + against kmeans_define_subtypes's heuristic, and optionally compute how + many further sub-clusters each existing finest-level label would be + apportioned if under-clustered major types were refined. + + Parameters + ---------- + obs : pd.DataFrame + Table containing the annotation columns, e.g. adata.obs. + level_keys : list of str + Column names in obs, ordered from coarsest (e.g. major cell type) + to finest. Must have at least 2 entries. + subcluster_size : int + Same meaning as in kmeans_define_subtypes. Used to compute the + expected label count per coarsest-level group. Must be >= 1. + min_frac : float + Same meaning as in kmeans_define_subtypes. Used only to flag + undersized labels below, never to drop them. + k_heuristic : callable + Same meaning as in kmeans_define_subtypes. Used to compute the + expected label count per coarsest-level group. + return_counts : bool + Compute and return the per-finest-label DataFrame described below. + When False, returns None and skips that computation entirely. + verbose : bool + Print a summary of finest-level label count, size range, and + expected label count for every coarsest-level group. + + Returns + ------- + pd.DataFrame or None + None unless return_counts is True. One row per + existing finest-level label, with columns level_keys, + n_cells (cells with that label), + major_type_fraction (n_cells as a fraction of its coarsest-level + group's total cells, the same quantity min_frac is compared + against), + current_k (number of finest labels in corresponding coarsest-level group), + expected_k (number of labels expected in that coarsest-level group, + based on k_heuristic and subcluster_size), and budget + (its Huntington-Hill-apportioned sub-cluster count if its major + type were refined, always >= 1, and equal to 1 throughout a + major type that isn't under-clustered). + + Raises + ------ + ValueError + If subcluster_size is < 1, or propagated from + check_subtype_consistency. + + Warns + ----- + UserWarning + Once per call, for each direction, naming every coarsest-level + group whose finest-level label count differs from what + kmeans_define_subtypes would produce by default for that group's + cell count (over-clustered and under-clustered reported + separately). Once per call for coarsest-level groups with any + finest-level label below min_frac of that group's cells, listing + each affected group as undersized-label-count/total-label-count + (see major_type_fraction, or return_counts, for which labels). + """ + if subcluster_size < 1: + raise ValueError(f"subcluster_size must be >= 1, got {subcluster_size}.") + + check_subtype_consistency(obs, level_keys) + + coarsest_key = level_keys[0] + finest_key = level_keys[-1] + + row_blocks = [] + under = [] + over = [] + small_summary = [] + for coarse_val, group in obs.groupby(coarsest_key, observed=True): + n_cells = len(group) + expected_k = _expected_k(n_cells, subcluster_size, k_heuristic) + + sizes = group.groupby(finest_key, observed=True).size() + n_labels = len(sizes) + + if n_labels > expected_k: + over.append(coarse_val) + if n_labels < expected_k: + under.append(coarse_val) + + small = sizes[sizes < min_frac * n_cells] + if len(small) > 0: + small_summary.append(f"{coarse_val}: {len(small)}/{n_labels}") + + if verbose: + print( + f" '{coarsest_key}'='{coarse_val}', '{finest_key}': " + f"{n_labels} labels (expected ~{expected_k}), " + f"sizes {int(sizes.min())}-{int(sizes.max())}, " + f"{n_cells} cells" + ) + + if return_counts: + budget_map = _huntington_hill_apportion(sizes.to_dict(), max(expected_k, n_labels)) + block = group[level_keys].drop_duplicates(subset=finest_key).reset_index(drop=True) + block["n_cells"] = block[finest_key].map(sizes).astype(int) + block["major_type_fraction"] = block["n_cells"] / n_cells + block["current_k"] = n_labels + block["expected_k"] = expected_k + block["budget"] = block[finest_key].map(budget_map).astype(int) + row_blocks.append(block) + + if under: + warnings.warn( + f"{len(under)} '{coarsest_key}' group(s) have fewer " + f"'{finest_key}' labels than subcluster_size/k_heuristic would " + f"produce by default: {under}.", + UserWarning, + stacklevel=find_stack_level(), + ) + if over: + warnings.warn( + f"{len(over)} '{coarsest_key}' group(s) have more " + f"'{finest_key}' labels than subcluster_size/k_heuristic would " + f"produce by default: {over}.", + UserWarning, + stacklevel=find_stack_level(), + ) + if small_summary: + warnings.warn( + f"{len(small_summary)} '{coarsest_key}' group(s) have " + f"'{finest_key}' labels below min_frac={min_frac}:\n " + + "\n ".join(small_summary), + UserWarning, + stacklevel=find_stack_level(), + ) + + if return_counts: + return pd.concat(row_blocks, ignore_index=True) + return None + + +def predefined_subtypes( + obs: pd.DataFrame, + level_keys: List[str], + subcluster_size: int = 10, + min_frac: float = 0.025, + k_heuristic: Callable[[int], float] = _default_k_heuristic, + verbose: bool = False, +) -> pd.DataFrame: + """ + Adapt a predefined, N-level subtype hierarchy after checking it is + consistent. + + Parameters + ---------- + obs : pd.DataFrame + Table containing the annotation columns, e.g. adata.obs. + level_keys : list of str + Column names in obs, ordered from coarsest (e.g. major cell type) + to finest. Must have at least 2 entries. + subcluster_size : int + Same meaning as in kmeans_define_subtypes. Used only to compute + the expected label count for the overclustering warning below. + Does not cap or drop labels. Must be >= 1. + min_frac : float + Same meaning as in kmeans_define_subtypes. Used only to flag + undersized labels below, never to drop them. + k_heuristic : callable + Same meaning as in kmeans_define_subtypes. Used only to compute + the expected label count for the overclustering warning below. + verbose : bool + Print a summary of finest-level label count, size range, and + expected label count for every coarsest-level group. + + Returns + ------- + pd.DataFrame + obs[level_keys], unchanged, with its original column names, order, + and index. + + Raises + ------ + ValueError + If subcluster_size is < 1, or propagated from + check_subtype_consistency if level_keys is invalid, a column is + missing from obs, contains missing values, or a label maps to + more than one label at a coarser level. + + Warns + ----- + UserWarning + Once, naming every coarsest-level group whose finest-level label + count exceeds what kmeans_define_subtypes would produce by + default for that group's cell count (and separately, once, for + groups with fewer labels than that). Once more for coarsest-level + groups with any finest-level label below min_frac of that group's + cells, listing each affected group's undersized-label-count out + of its total label count. + """ + plan_subtype_refinement( + obs, level_keys, + subcluster_size=subcluster_size, + min_frac=min_frac, + k_heuristic=k_heuristic, + return_counts=False, + verbose=verbose, + ) + return obs[level_keys].copy() + - cluster_ids = [] - for sc_label in np.unique(labels_valid): - sc_mask = labels_valid == sc_label +def refine_predefined_subtypes( + X, + obs: pd.DataFrame, + level_keys: List[str], + plan: pd.DataFrame, + min_frac: float = 0.04, + th_inner_logfold: float = 0.75, + max_genes: int = 500, + max_cells_per_label: int = 10000, + name_sep: str = "_", + random_state: int = 42, + verbose: bool = False, +) -> pd.DataFrame: + """ + Further sub-cluster the existing finest-level labels of under-clustered + major types, using the budget from plan_subtype_refinement. + + For each major type with at least one label whose budget exceeds 1, + genes are narrowed once for the whole major type (comparing it against + every other major type, the same comparison kmeans_define_subtypes + uses for its own top-level clustering) and that gene space is shared + by every label being refined within it. Labels with budget == 1 + (including every label of a major type with none flagged) are passed + through unchanged, aside from the naming suffix. + + Parameters + ---------- + X : array-like or sparse matrix + Cell x gene expression matrix, row-aligned with obs. + obs : pd.DataFrame + Table containing the annotation columns, e.g. predefined_subtypes's + output. + level_keys : list of str + Column names in obs, ordered from coarsest to finest. Must match + the level_keys plan was built from. + plan : pd.DataFrame + plan_subtype_refinement's output with return_counts=True: one row + per existing finest-level label, including its budget. + min_frac : float + Minimum fraction of a label's sampled cells (after max_cells_per_label) + a new sub-cluster must contain to be kept; smaller ones are dropped + as noise. Never applied below an absolute floor of 10 cells, + regardless of how small min_frac * sampled cells comes out. + th_inner_logfold : float + Log-fold threshold for narrowing to marker-like genes before + clustering. Same meaning as in kmeans_define_subtypes. + max_genes : int + Gene narrowing before kmeans clustering only runs when X has more + than this many genes, and keeps at most this many when it does. + max_cells_per_label : int + Cap on how many of a label's cells are used for clustering; labels + larger than this are randomly subsampled. Cells left out are not + assigned a subtype. + name_sep : str + Separator used to build new subtype names, e.g. + f"{old_label}{name_sep}{i}". + random_state : int + Random seed for reproducibility. + verbose : bool + Print progress messages. - if issparse(X_ct_valid): - sc_centroid = np.asarray(X_ct_valid[sc_mask].mean(axis=0)).flatten() - else: - sc_centroid = X_ct_valid[sc_mask].mean(axis=0) + Returns + ------- + pd.DataFrame + obs[level_keys] with an additional column, level_keys[-1] + + "_kmeans", holding the refined labels. + + Warns + ----- + UserWarning + If any candidate sub-cluster was dropped as noise, or if every + candidate sub-cluster of a label was dropped (a single fallback + subtype is assigned instead). + """ + np.random.seed(random_state) - sub_centroids_list.append(sc_centroid) - cluster_ids.append(len(sub_centroids_list) - 1) + coarsest_key = level_keys[0] + finest_key = level_keys[-1] + n_genes = X.shape[1] - clusters[ct] = cluster_ids + coarsest_values = obs[coarsest_key].values + finest_values = obs[finest_key].values + major_types = np.unique(coarsest_values) + + major_centroids = None + major_ratios = None + if plan["budget"].gt(1).any() and n_genes > max_genes: + major_centroids, major_ratios = _capped_group_centroids(X, coarsest_values) + if verbose: + print(f"Preparing to refine {plan['budget'].gt(1).sum()} '{finest_key}' labels ") + + labels_out = np.full(X.shape[0], None, dtype=object) + + for ct_pos, mt in enumerate(major_types): + group_plan = plan[plan[coarsest_key] == mt] + + # if no subclustering required -> passthrough + if (group_plan["budget"] <= 1).all(): + for _, row in group_plan.iterrows(): + label = row[finest_key] + labels_out[finest_values == label] = f"{label}{name_sep}0" + continue + + # Narrow down to marker-like genes once for the whole major type. + gene_indices = np.arange(n_genes) + if n_genes > max_genes: + gene_indices = _select_kmeans_genes( + ct_centroid=major_centroids[ct_pos], + major_centroids=major_centroids, + major_ratios=major_ratios, + ct_name=mt, + th_logfold=th_inner_logfold, + max_genes=max_genes, + ) + + # for each finest-level label, sub-cluster it into its budgeted number of subtypes + for _, row in group_plan.iterrows(): + label = row[finest_key] + budget = row["budget"] + label_indices = np.where(finest_values == label)[0] + + # Skip clustering for labels with budget 1, just rename them with the suffix. + if budget == 1: + labels_out[label_indices] = f"{label}{name_sep}0" + continue + + # Bound clustering cost; cells left out simply get no subtype. + sampled_indices = label_indices + if len(label_indices) > max_cells_per_label: + sampled_indices = np.random.choice(label_indices, max_cells_per_label, replace=False) + + X_label = X[sampled_indices] + + if verbose: + print(f" Clustering {len(sampled_indices)} '{label}' cells into " + f"~{budget} subtypes (using {len(gene_indices)} genes)...") + + sub_labels = _kmeans_subcluster( + X_ct=X_label, + gene_indices=gene_indices, + K=budget, + min_cells=max(min_frac * len(sampled_indices), 10), + random_state=random_state, + ) + + valid_mask = sub_labels >= 0 + n_dropped_cells = int((~valid_mask).sum()) + + # Fall back to one subtype for the whole label (not just the + # sampled cells) if every candidate sub-cluster was too small + # to keep -- no clustering distinction survives to preserve, + # so including cells left out by max_cells_per_label only + # makes the resulting single centroid less noisy. + if valid_mask.sum() == 0: + labels_out[label_indices] = f"{label}{name_sep}0" + warnings.warn( + f"All candidate subtypes of '{label}' were below " + f"min_frac={min_frac} and dropped as noise; assigned a " + f"single subtype '{label}{name_sep}0' instead.", + UserWarning, + stacklevel=find_stack_level(), + ) + continue + + # Name and assign each surviving sub-cluster; noise cells (-1) + # are skipped, leaving them None in labels_out. + unique_sub_labels = np.unique(sub_labels[valid_mask]) + for i, sc_label in enumerate(unique_sub_labels): + cell_mask = sub_labels == sc_label + labels_out[sampled_indices[cell_mask]] = f"{label}{name_sep}{i}" + + if n_dropped_cells > 0: + warnings.warn( + f"Dropped {n_dropped_cells} cell(s) from '{label}' as " + f"noise (subtype(s) below min_frac={min_frac} of its " + f"{len(sampled_indices)} sampled cells).", + UserWarning, + stacklevel=find_stack_level(), + ) + + result = obs[level_keys].copy() + result[f"{finest_key}_kmeans"] = pd.Series(labels_out, index=obs.index, dtype=object) + return result + + +def summarize_subtypes(X, labels: pd.DataFrame, verbose: bool = False) -> Dict: + """ + Aggregate cells into per-subtype centroids from a subtype hierarchy. + + Parameters + ---------- + X : array-like or sparse matrix + Cell x gene expression matrix, row-aligned with labels. + labels : pd.DataFrame + Subtype hierarchy for each cell, as returned by + kmeans_define_subtypes or predefined_subtypes: one column per + level, ordered from coarsest to finest. Cells with a missing + finest-level label are excluded. + verbose : bool + Print a summary of how many cells and groups were aggregated. + + Returns + ------- + dict + 'centroids' : mean expression profile per surviving finest-level + group. + 'hierarchy' : labels' columns for each surviving group, aligned + with 'centroids' rows. + 'ratios' : fraction of cells per coarsest-level group, restricted + to groups with at least one surviving cell. + + Raises + ------ + ValueError + Propagated from check_subtype_consistency if labels' columns are + inconsistent among cells with a finest-level label. + + Warns + ----- + UserWarning + For any finest-level group with fewer than 3 cells. + """ + coarsest_key = labels.columns[0] + finest_key = labels.columns[-1] + + valid_mask = labels[finest_key].notna().to_numpy() + X_valid = X[valid_mask] + labels_valid = labels[valid_mask].reset_index(drop=True) + + check_subtype_consistency(labels_valid, list(labels_valid.columns)) + + centroids_list = [] + hierarchy_rows = [] + for label, group in labels_valid.groupby(finest_key, observed=True): + idx = group.index.to_numpy() + + if len(idx) < 3: + warnings.warn( + f"'{finest_key}'='{label}' has only {len(idx)} cell(s); " + f"centroids from fewer than 3 cells may be unreliable.", + UserWarning, + stacklevel=find_stack_level(), + ) + + if issparse(X_valid): + centroid = np.asarray(X_valid[idx].mean(axis=0)).flatten() + else: + centroid = X_valid[idx].mean(axis=0) + centroids_list.append(centroid) + hierarchy_rows.append(group.iloc[0,:]) if issparse(X): - sub_centroids = vstack([csr_matrix(c) for c in sub_centroids_list]) + centroids = vstack([csr_matrix(c) for c in centroids_list]) + else: + centroids = np.array(centroids_list) + + hierarchy = pd.DataFrame(hierarchy_rows).reset_index(drop=True) + + # Ratios use each coarsest-level group's full original cell count (not + # reduced by finest-level drops within it), restricted to groups with + # at least one surviving cell. + full_counts = labels[coarsest_key].value_counts() + surviving = full_counts[full_counts.index.isin(labels_valid[coarsest_key].unique())] + ratios = (surviving / surviving.sum()).to_dict() + + if verbose: + print( + f"Summarized {len(labels_valid)} cells into " + f"{len(centroids_list)} '{finest_key}' groups across " + f"{len(ratios)} '{coarsest_key}' groups." + ) + + return {"centroids": centroids, "hierarchy": hierarchy, "ratios": ratios} + + +def _de_scores_logfc_rank(centroids: np.ndarray) -> np.ndarray: + """R-style median-rank log-fold-change gene scores; lower is more distinctive.""" + C, G = centroids.shape + centroids = centroids + 1e-9 + + gene_scores = np.zeros((C, G), dtype=np.float32) + for i in range(C): + this_ct = np.tile(centroids[i:i + 1, :], (C - 1, 1)) + other_ct = np.delete(centroids, i, axis=0) + logfc = np.log(this_ct / other_ct) + + ranks = np.empty_like(logfc) + for j in range(C - 1): + ranks[j, :] = G - np.argsort(np.argsort(logfc[j, :])) + + gene_scores[i, :] = np.median(ranks, axis=0) + + return gene_scores + + +_DE_SCORE_METHODS: Dict[str, Callable[[np.ndarray], np.ndarray]] = { + "logfc_rank": _de_scores_logfc_rank, +} + + +def select_de_genes( + centroids, + max_genes: int = 5000, + method: Union[str, Callable[[np.ndarray], np.ndarray]] = "logfc_rank", + return_scores: bool = False, + verbose: bool = False, +): + """ + Select the most differentially expressed genes across subtype centroids. + + Parameters + ---------- + centroids : array-like or sparse matrix + Subtype centroids, shape (C, G): one row per subtype, one column + per gene. + max_genes : int + Maximum number of genes to select. + method : str or callable + The name of a built-in method, or a custom scoring callable. + 'logfc_rank' (the default R-style method) is currently + the only built-in method implemented. A callable takes the dense + (C, G) centroids and returns a (C, G) float array of per-subtype, + per-gene scores, where lower means more distinctive. It must not + mutate its input. + return_scores : bool + Also return the (C, G) score matrix used for selection. + verbose : bool + Print progress messages. + + Returns + ------- + np.ndarray + Indices of the selected genes. + np.ndarray, optional + The (C, G) score matrix used for selection, if return_scores is + True. None if max_genes already covers every gene, since no + scoring took place. + + Raises + ------ + ValueError + If method is a string not found in the scoring method registry. + """ + n_genes = centroids.shape[1] + + if n_genes <= max_genes: + top_genes = np.arange(n_genes) + return (top_genes, None) if return_scores else top_genes + + if isinstance(method, str): + if method not in _DE_SCORE_METHODS: + raise ValueError(f"Unknown method '{method}'. Available: {list(_DE_SCORE_METHODS)}.") + scorer = _DE_SCORE_METHODS[method] else: - sub_centroids = np.array(sub_centroids_list) + scorer = method - surviving = {ct: major_ratios[ct] for ct in clusters.keys() if ct in major_ratios} - total = sum(surviving.values()) - if total > 0: - major_ratios = {ct: v / total for ct, v in surviving.items()} + centroids_dense = centroids.toarray() if issparse(centroids) else np.asarray(centroids) if verbose: - print(f"Created {len(sub_centroids_list)} sub-clusters from {len(clusters)} cell types") + method_name = method if isinstance(method, str) else getattr(method, "__name__", "custom") + print( + f"Computing gene scores ({method_name}) for {centroids_dense.shape[0]} " + f"clusters, {n_genes} genes..." + ) + + gene_scores = scorer(centroids_dense) + top_genes = np.argsort(gene_scores.min(axis=0))[:max_genes] + + if verbose: + print(f"Selected {len(top_genes)} DE genes") + + return (top_genes, gene_scores) if return_scores else top_genes + + +def _warn_mt_genes(adata: AnnData, context: str) -> None: + """ + Warn (with per-prefix gene counts) if and how many MT-/HLA-/RPL-prefixed genes are + present in the AnnData. These genes are not removed automatically. + """ + var_names = adata.var_names + n_mt = int(var_names.str.startswith('MT-').sum()) + n_hla = int(var_names.str.startswith('HLA-').sum()) + n_rpl = int(var_names.str.startswith('RPL').sum()) + if n_mt or n_hla or n_rpl: + warnings.warn( + f"Data contains {n_mt} MT-, {n_hla} HLA-, {n_rpl} RPL-prefixed " + f"gene(s) out of {adata.shape[1]} total. They are not removed " + f"automatically -- filter them yourself first if they might bias " + f"{context} (pass warn_mt=False to silence this message).", + UserWarning, + stacklevel=find_stack_level(), + ) + + +def validate_reference_input( + adata: AnnData, + cell_type_key: Union[str, List[str]], + max_input_genes: int, + warn_mt: bool, + check_counts: bool, + check_nonnegative: bool, +) -> None: + """ + Validate that adata is ready for reference aggregation, without modifying + adata. + + QC filtering, MT/HLA/RPL removal, and HVG selection are left to the caller + before this function is invoked. Called automatically by setup_reference(), + call it directly yourself if you are running the individual preprocessing + steps rather than the wrapper. + + Parameters + ---------- + adata : AnnData + Reference single-cell data to validate. + cell_type_key : str or list of str + Key(s) in adata.obs expected to contain cell type annotations. A + single key or a list of keys are both checked the same way, every + one of them. + max_input_genes : int + Raise if adata has more genes than this. + warn_mt : bool + Warn (not raise) if MT-/HLA-/RPL-prefixed genes are present. + check_counts : bool + Raise if adata.X does not look like raw (integer-valued) counts. + check_nonnegative : bool + Raise if adata.X contains negative values. + + Raises + ------ + ValueError + If any of the checks above fail. + """ + cell_type_keys = [cell_type_key] if isinstance(cell_type_key, str) else cell_type_key + for key in cell_type_keys: + if key not in adata.obs.columns: + raise ValueError( + f"cell_type_key '{key}' not found in adata.obs. " + f"Available columns: {list(adata.obs.columns)}" + ) + if adata.obs[key].isna().any(): + n_na = int(adata.obs[key].isna().sum()) + raise ValueError( + f"cell_type_key '{key}' has {n_na} missing (NaN) value(s). " + "Remove or annotate these cells before calling setup_reference()." + ) + + X = adata.X + data = X.data if issparse(X) else np.asarray(X).ravel() + if check_nonnegative and data.size and (data < 0).any(): + raise ValueError( + "adata.X contains negative values. setup_reference() expects raw " + "(non-negative) counts, not scaled/z-scored data. If this is " + "expected (e.g. scaled spatial proteomics data), set " + "check_nonnegative=False. " + ) + if check_counts and data.size and not np.allclose(data, np.round(data)): + raise ValueError( + "adata.X does not look like raw counts (non-integer values found). " + "setup_reference() expects raw counts by default. If this is expected " + "(e.g. background-corrected data), pass check_counts=False." + ) + + empty_cells = int((np.asarray(X.sum(axis=1)).ravel() == 0).sum()) + if empty_cells > 0: + raise ValueError( + f"{empty_cells} cell(s) have zero total counts. Run " + "sc.pp.filter_cells(adata, min_counts=1) (or similar) before calling " + "setup_reference()." + ) + empty_genes = int((np.asarray(X.sum(axis=0)).ravel() == 0).sum()) + if empty_genes > 0: + raise ValueError( + f"{empty_genes} gene(s) have zero total counts across all cells. Run " + "sc.pp.filter_genes(adata, min_cells=1) (or similar) before calling " + "setup_reference()." + ) - return { - 'major_centroids': major_centroids, - 'major_ratios': major_ratios, - 'sub_centroids': sub_centroids, - 'clusters': clusters - } + if warn_mt: + _warn_mt_genes(adata, context="sub-clustering") + + if adata.shape[1] > max_input_genes: + raise ValueError( + f"adata has {adata.shape[1]} genes, exceeding max_input_genes=" + f"{max_input_genes}. Sub-clustering and DE-gene ranking do not scale " + "well past this many input genes -- subset genes yourself first (e.g. " + "sc.pp.highly_variable_genes) or raise max_input_genes if you accept " + "the extra compute cost." + ) def setup_reference( adata: AnnData, - cell_type_key: str, + cell_type_key: Union[str, List[str]], subcluster_size: int = 10, max_genes: int = 5000, - remove_mt: bool = True, + max_input_genes: int = 5000, th_inner_logfold: float = 0.75, + min_frac: float = 0.025, + refine_undersized: bool = False, random_state: int = 42, + warn_mt: bool = True, + check_counts: bool = True, + check_nonnegative: bool = True, + kws_kmeans: Optional[Dict] = None, + kws_predefined: Optional[Dict] = None, + kws_refine: Optional[Dict] = None, + kws_de_genes: Optional[Dict] = None, verbose: bool = False, - copy: bool = True -) -> Dict: +) -> AnnData: """ - Process reference single-cell RNA-seq data for DOT (R-compatible version). + Aggregate and sub-cluster a reference single-cell RNA-seq dataset for DOT + (R-compatible algorithm). + + This does not run QC filtering, MT/HLA/RPL removal, or HVG selection -- + those are standard scanpy steps left to the caller. adata is expected to + already be filtered (no empty cells/genes), with raw counts in .X (for single-cell RNA-seq data). Parameters ---------- adata : AnnData - Reference single-cell data with raw counts in .X - cell_type_key : str - Key in adata.obs containing cell type annotations + Reference single-cell data, pre-filtered, with raw (positive) counts in .X. + cell_type_key : str or list of str + Column(s) in adata.obs with the cell type hierarchy, ordered + coarsest to finest. A single column (as a str, or a list with one + entry) means subtypes are discovered with k-means; a list with two + or more entries means subtypes are taken as given via + predefined_subtypes. subcluster_size : int - Maximum number of sub-clusters per cell type + Maximum number of subtypes per cell type. Caps clustering directly + on the k-means path; used as the reference cap for the + overclustering warning on the predefined path. max_genes : int - Maximum number of genes to use - remove_mt : bool - Whether to remove mitochondrial / ribosomal genes + Maximum number of DE genes to keep in the output. + max_input_genes : int + Maximum number of genes adata is allowed to have on input; raises if + exceeded, since sub-clustering/DE-ranking cost grows with gene count. + Independent of max_genes -- run your own HVG selection first if you + need to go above this. th_inner_logfold : float - Log-fold threshold for gene selection in sub-clustering + Log-fold threshold for gene selection in sub-clustering (k-means + path only). + min_frac : float + Minimum fraction of a cell type's cells a subtype must contain; + smaller ones are dropped as noise on the k-means path, or flagged + (never dropped) on the predefined path. 0 disables k-means dropping. + refine_undersized : bool + Predefined path only. If True, major types whose finest-level + label count falls short of subcluster_size/k_heuristic's expected + count are further sub-clustered with k-means, using budgets from + plan_subtype_refinement. random_state : int - Random seed for reproducibility + Random seed for reproducibility (k-means path only). + warn_mt : bool + Warn (with per-prefix gene counts) if MT-/HLA-/RPL-prefixed genes are + present in the AnnData. They are not removed automatically. + check_counts : bool + Check that adata.X looks like raw (integer-valued) counts; set this + False to allow non-integer values, e.g. for background-corrected data. + check_nonnegative : bool + Check that adata.X has no negative values; set this False for + inherently scaled/centered data, e.g. spatial proteomics. Note this + only lifts the input check -- it does not change the sub-clustering + gene-selection or DE-ranking math (_select_kmeans_genes, + select_de_genes), which compute log-fold-changes and assume + non-negative values. With negative input those can silently produce + NaNs (via log of a non-positive ratio) rather than raising, so treat + this as a known limitation, not a validated code path, until that + math is revisited. + kws_kmeans : dict, optional + Extra keyword arguments forwarded to kmeans_define_subtypes() + (e.g. k_heuristic, max_cells_per_type, name_sep); see its + docstring. Only used on the k-means path. + kws_predefined : dict, optional + Extra keyword arguments forwarded to plan_subtype_refinement() + (e.g. k_heuristic); see its docstring. Only used on the predefined + path. + kws_refine : dict, optional + Extra keyword arguments forwarded to refine_predefined_subtypes() + (e.g. min_frac, max_genes, max_cells_per_label, name_sep); see its + docstring. Only used when refine_undersized is True. Its + max_cells_per_label is analogous to kmeans_define_subtypes's + max_cells_per_type, but scoped to each existing finest-level label + rather than each major type. min_frac is set to 0.04 by default, + and requires at least 10 cells per resulting candidate sub-cluster, + which is stricter than the default 0.025 min_frac on the k-means. + kws_de_genes : dict, optional + Extra keyword arguments forwarded to select_de_genes() (e.g. + method, return_scores); see its docstring. Used on both paths. verbose : bool - Print progress messages - copy : bool - Whether to copy adata before processing + Print progress messages. Returns ------- - dict - 'X_sparse', 'clusters', 'ratios', 'genes' + AnnData + One row per subtype centroid. X holds the centroid x DE-gene + matrix, obs holds the subtype hierarchy (one column per level, + coarsest to finest), var_names holds the DE genes, uns['ratios'] + holds each coarsest-level group's cell fraction, and + uns['level_keys'] holds obs's hierarchy column names in the same + coarsest-to-finest order (an explicit record of the order, since + obs itself may later gain unrelated columns from other tools). + On the predefined path (cell_type_key has 2+ entries), + uns['subtype_plan'] holds plan_subtype_refinement's per-label + DataFrame. When refine_undersized is also True, it additionally + has an n_subtypes_made column recording how many sub-clusters + each existing label actually produced (as opposed to budget, its + target before clustering was attempted) -- absent when + refine_undersized is False, since nothing was actually + clustered. Not present on the k-means path. """ if verbose: print("=" * 60) @@ -349,80 +1283,97 @@ def setup_reference( print("=" * 60) print(f"Input shape: {adata.shape}") - if copy: - adata = adata.copy() - - if verbose: - print("\nRunning basic QC...") - sc.pp.filter_cells(adata, min_counts=1) - sc.pp.filter_genes(adata, min_cells=1) - - if remove_mt: - mt_mask = adata.var_names.str.startswith(('MT-', 'HLA-', 'RPL')) - n_mt = mt_mask.sum() - if n_mt > 0: - adata = adata[:, ~mt_mask].copy() - if verbose: - print(f"Removed {n_mt} MT/HLA/RPL genes") + kws_kmeans = kws_kmeans or {} + kws_predefined = kws_predefined or {} + kws_refine = kws_refine or {} + kws_de_genes = kws_de_genes or {} + + keys = [cell_type_key] if isinstance(cell_type_key, str) else list(cell_type_key) + coarsest_key = keys[0] + + validate_reference_input( + adata, + cell_type_key=keys, + max_input_genes=max_input_genes, + warn_mt=warn_mt, + check_counts=check_counts, + check_nonnegative=check_nonnegative, + ) X = adata.X - annotations = adata.obs[cell_type_key].values.astype(str) genes = adata.var_names.values - vg_genes = max(5000, max_genes) - if adata.shape[1] > vg_genes: - if verbose: - print(f"\nSelecting {vg_genes} highly variable genes...") - - adata_hvg = adata.copy() - adata_hvg.layers['counts'] = adata.X.copy() - - sc.pp.highly_variable_genes( - adata_hvg, - n_top_genes=vg_genes, - flavor='seurat_v3', - layer='counts', - subset=False + if verbose: + print("\nDefining subtypes...") + + plan = None + if len(keys) == 1: + annotations = adata.obs[coarsest_key].values.astype(str) + labels = kmeans_define_subtypes( + X, annotations, + subcluster_size=subcluster_size, + th_inner_logfold=th_inner_logfold, + min_frac=min_frac, + random_state=random_state, + verbose=verbose, + annotations_key=coarsest_key, + **kws_kmeans, ) + else: + plan = plan_subtype_refinement( + adata.obs, keys, + subcluster_size=subcluster_size, + min_frac=min_frac, + return_counts=True, + verbose=verbose, + **kws_predefined, + ) + if refine_undersized: + labels = refine_predefined_subtypes( + X, adata.obs[keys].copy(), keys, plan, + th_inner_logfold=th_inner_logfold, + random_state=random_state, + verbose=verbose, + **kws_refine, + ) + finest_key = keys[-1] + n_subtypes_made = labels.groupby(finest_key)[f"{finest_key}_kmeans"].nunique() + plan["n_subtypes_made"] = plan[finest_key].map(n_subtypes_made).astype(int) - hvg_genes = adata_hvg.var_names[adata_hvg.var['highly_variable']].tolist() - adata = adata[:, hvg_genes].copy() - X = adata.X - genes = adata.var_names.values + if verbose: + n_before, n_after = len(plan), int(plan["n_subtypes_made"].sum()) + print(f"\nAdded {n_after - n_before} subclusters to finest cell " + f"type levels ({n_before} -> {n_after}).") + else: + labels = adata.obs[keys].copy() if verbose: - print(f"\nAfter filtering: {X.shape}") - print("\nAggregating and sub-clustering cell types...") - - ref_agg = _aggregate_reference( - X=X, - annotations=annotations, - cluster_size=subcluster_size, - th_inner_logfold=th_inner_logfold, - random_state=random_state, - verbose=verbose - ) + print("\nSummarizing subtypes into centroids...") - if verbose: - print("\nSelecting differentially expressed genes (R-style)...") + summary = summarize_subtypes(X, labels, verbose=verbose) - de_genes = _get_de_genes_r_style( - centroids=ref_agg['sub_centroids'], - max_genes=max_genes, - verbose=verbose - ) + if verbose: + print("\nSelecting differentially expressed genes...") - if issparse(ref_agg['sub_centroids']): - X_subset = ref_agg['sub_centroids'][:, de_genes] + de_result = select_de_genes(summary['centroids'], max_genes=max_genes, verbose=verbose, **kws_de_genes) + if kws_de_genes.get('return_scores'): + de_genes, de_scores = de_result else: - X_subset = ref_agg['sub_centroids'][:, de_genes] - - result = { - 'X_sparse': X_subset, - 'clusters': ref_agg['clusters'], - 'ratios': ref_agg['major_ratios'], - 'genes': genes[de_genes] - } + de_genes, de_scores = de_result, None + X_subset = summary['centroids'][:, de_genes] + + obs = summary['hierarchy'] + obs.index = summary['hierarchy'][summary['hierarchy'].columns[-1]] + obs.index.name = None + + var = pd.DataFrame(index=genes[de_genes]) + ref_adata = AnnData(X=X_subset, obs=obs, var=var) + if de_scores is not None: + ref_adata.uns['de_scores'] = de_scores + ref_adata.uns['ratios'] = summary['ratios'] + ref_adata.uns['level_keys'] = list(obs.columns) + if plan is not None: + ref_adata.uns['subtype_plan'] = plan if verbose: print(f"\n{'=' * 60}") @@ -434,7 +1385,38 @@ def setup_reference( print(f" - Sparsity: {sparsity:.2%}") print(f"{'=' * 60}") - return result + return ref_adata + + +def _validate_spatial_input( + adata: AnnData, + warn_mt: bool, +) -> None: + """ + Validate/warn on spatial input before processing. Mirrors + validate_reference_input() + """ + if warn_mt: + _warn_mt_genes(adata, context="the spatial similarity computation") + + +def _symmetric_pairs_matrix( + i_idx: np.ndarray, + j_idx: np.ndarray, + weights: np.ndarray, + n: int, +) -> csr_matrix: + """ + Build a symmetric spot-by-spot sparse matrix from a one-directional + edge list (each pair listed once, with i_idx < j_idx). + """ + return csr_matrix( + ( + np.concatenate([weights, weights]), + (np.concatenate([i_idx, j_idx]), np.concatenate([j_idx, i_idx])), + ), + shape=(n, n), + ) def setup_spatial( @@ -444,11 +1426,11 @@ def setup_spatial( th_nonspatial: float = 0.0, th_gene_low: float = 0.01, th_gene_high: float = 0.99, - remove_mt: bool = True, + warn_mt: bool = True, radius: str = 'auto', verbose: bool = False, copy: bool = True -) -> Dict: +) -> AnnData: """ Process spatial transcriptomics data. @@ -466,8 +1448,9 @@ def setup_spatial( Minimum fraction of spots a gene must be expressed in th_gene_high : float Maximum fraction of spots a gene can be expressed in - remove_mt : bool - Remove MT/RPL genes + warn_mt : bool + Warn (with per-prefix gene counts) if MT-/HLA-/RPL-prefixed genes are + present in the AnnData. They are not removed automatically. radius : str or float Spatial radius ('auto' or numeric value) verbose : bool @@ -477,8 +1460,16 @@ def setup_spatial( Returns ------- - dict - 'X_sparse', 'coords', 'genes', 'pairs' (if th_spatial > 0) + AnnData + Same spots as adata (X and var_names reflect gene-frequency + filtering; obs is preserved from the input). obsm['spatial'] holds + the 2D coordinates. obsp['dot_spatial_pairs'] holds the + spot-neighbor graph as a symmetric sparse matrix of + cosine-similarity edge weights, present only if any pairs were + found. Not named obsp['spatial_connectivities'] deliberately -- + unlike squidpy's slot of that name, these weights come from + gene-expression similarity within a spatial radius, not from + spatial distance alone. """ if copy: adata = adata.copy() @@ -488,18 +1479,12 @@ def setup_spatial( coords = np.asarray(adata.obsm[spatial_key]) if coords.shape[1] > 2: coords = coords[:, :2] + adata.obsm['spatial'] = coords if verbose: print(f"Processing spatial data: {adata.shape}") - # Remove MT genes - if remove_mt: - mt_mask = adata.var_names.str.startswith(('MT-', 'HLA-', 'RPL')) - n_mt = mt_mask.sum() - if n_mt > 0: - adata = adata[:, ~mt_mask].copy() - if verbose: - print(f"Removed {n_mt} MT/HLA/RPL genes") + _validate_spatial_input(adata, warn_mt=warn_mt) # Gene frequency filter if th_gene_high < 1 or th_gene_low > 0: @@ -510,13 +1495,6 @@ def setup_spatial( print(f"Filtered to {adata.shape[1]} genes") X = adata.X - genes = adata.var_names.values - - result = { - 'X_sparse': X, - 'coords': coords, - 'genes': genes - } # Spatial pairs if th_spatial > 0: @@ -634,12 +1612,9 @@ def setup_spatial( if verbose: print(f"Found {len(i_idx)} spatial pairs") - # Return as dict - result['pairs'] = { - 'i': i_idx, - 'j': j_idx, - 'w': weights - } + adata.obsp['dot_spatial_pairs'] = _symmetric_pairs_matrix( + i_idx, j_idx, weights, n=coords.shape[0] + ) elif verbose: print("No spatial pairs found") elif verbose: @@ -650,4 +1625,4 @@ def setup_spatial( if issparse(X): print(f"Sparsity: {1 - X.nnz / (X.shape[0] * X.shape[1]):.2%}") - return result \ No newline at end of file + return adata \ No newline at end of file diff --git a/dotpy/visualization.py b/dotpy/visualization.py index fc8f6e1..3a132d3 100644 --- a/dotpy/visualization.py +++ b/dotpy/visualization.py @@ -5,14 +5,15 @@ """ import numpy as np +import pandas as pd import matplotlib.pyplot as plt from matplotlib.colors import Normalize -from typing import Optional, Tuple +from typing import Optional, Tuple, Union def plot_spatial_weights( coords: np.ndarray, - weights: np.ndarray, + weights: Union[np.ndarray, pd.DataFrame], cell_types: Optional[list] = None, normalize: bool = True, ncols: int = 4, @@ -28,6 +29,7 @@ def plot_spatial_weights( dpi: int = 150 ) -> plt.Figure: """Plot spatial distribution of cell type weights.""" + weights = np.asarray(weights) if normalize: row_sums = weights.sum(axis=1, keepdims=True) row_sums[row_sums == 0] = 1 @@ -86,7 +88,7 @@ def plot_spatial_weights( def plot_cell_type_proportions( - weights: np.ndarray, + weights: Union[np.ndarray, pd.DataFrame], cell_types: Optional[list] = None, figsize: Tuple[float, float] = (10, 6), colors: Optional[list] = None, @@ -94,6 +96,7 @@ def plot_cell_type_proportions( dpi: int = 150 ) -> plt.Figure: """Plot overall cell type proportions across all spots.""" + weights = np.asarray(weights) n_ct = weights.shape[1] if cell_types is None: cell_types = [f"CT{i+1}" for i in range(n_ct)] diff --git a/example.py b/example.py index e13b93a..3b50255 100644 --- a/example.py +++ b/example.py @@ -100,14 +100,19 @@ def example_workflow(): # 3. Run optimisation print(f"\n-- Optimisation (batch_size={batch}, device={device}) --") - dot = DOT(sp, ref, ls_solution=True, batch_size=batch, device=device) + dot = DOT( + sp, ref, + mode='highres', + ratios_weight=0.0, + ls_solution=True, + batch_size=batch, + device=device, + ) ckpt_dir = './checkpoints_example' Path(ckpt_dir).mkdir(exist_ok=True) dot.fit( - mode='highres', - ratios_weight=0.0, iterations=30, gap_threshold=0.01, verbose=True, @@ -120,8 +125,8 @@ def example_workflow(): weights = dot.get_weights(normalize=True) cts = dot.get_cell_types() print(f"\nWeights: {weights.shape} | Cell types: {cts}") - for i, ct in enumerate(cts): - print(f" {ct}: mean={weights[:, i].mean():.4f}") + for ct in cts: + print(f" {ct}: mean={weights[ct].mean():.4f}") # 5. Visualize print("\n-- Creating visualizations --") @@ -159,11 +164,10 @@ def example_resume(): ref = setup_reference(_make_reference(), cell_type_key='cell_type', verbose=False) sp = setup_spatial(_make_spatial(), verbose=False) - dot = DOT(sp, ref, batch_size=500, device=device) + dot = DOT(sp, ref, mode='highres', batch_size=500, device=device) print(f"Resuming optimization on {device}...") dot.fit( - mode='highres', iterations=50, resume_from=ckpt, checkpoint_dir='./checkpoints_example', @@ -202,10 +206,9 @@ def example_memory_constrained(): verbose=True ) - dot = DOT(sp, ref, batch_size=100, device=device) # Small batch size + dot = DOT(sp, ref, mode='highres', batch_size=100, device=device) # Small batch size dot.fit( - mode='highres', iterations=20, verbose=True, use_mixed_precision=True, # Use float16 to save memory @@ -237,11 +240,10 @@ def example_high_quality(): verbose=True ) - dot = DOT(sp, ref, batch_size=500, device=device) + dot = DOT(sp, ref, mode='highres', batch_size=500, device=device) print(f"Running high-quality optimization on {device}...") dot.fit( - mode='highres', iterations=200, # More iterations gap_threshold=0.001, # Tighter convergence verbose=True, @@ -265,8 +267,8 @@ def example_visualization(): ref = setup_reference(_make_reference(), cell_type_key='cell_type', verbose=False) sp = setup_spatial(_make_spatial(), verbose=False) - dot = DOT(sp, ref, device=device) - dot.fit(mode='highres', iterations=30, verbose=False) + dot = DOT(sp, ref, mode='highres', device=device) + dot.fit(iterations=30, verbose=False) weights = dot.get_weights(normalize=True) cts = dot.get_cell_types() diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 0000000..7d2b13f --- /dev/null +++ b/pytest.ini @@ -0,0 +1,3 @@ +[pytest] +testpaths = tests +filterwarnings = default diff --git a/requirements-test.txt b/requirements-test.txt new file mode 100644 index 0000000..e079f8a --- /dev/null +++ b/requirements-test.txt @@ -0,0 +1 @@ +pytest diff --git a/run_dot_cli.py b/run_dot_cli.py index 1dccf3d..acbd3a5 100644 --- a/run_dot_cli.py +++ b/run_dot_cli.py @@ -261,6 +261,8 @@ def main(): dot = DOT( spatial_processed, ref_processed, + mode=args.mode, + ratios_weight=args.ratios_weight, batch_size=args.batch_size, device=device ) @@ -271,8 +273,6 @@ def main(): ckpt_dir = str(Path(args.checkpoint_dir) / str(sample_id)) dot.fit( - mode=args.mode, - ratios_weight=args.ratios_weight, iterations=args.iterations, verbose=args.verbose, use_mixed_precision=args.mixed_precision, @@ -281,17 +281,11 @@ def main(): resume_from=args.resume_from, ) - # Get results - weights = dot.get_weights(normalize=True) + # Get results -- get_weights() already returns a DataFrame indexed + # by spot and labeled by cell type. + weights_df = dot.get_weights(normalize=True) cell_types = dot.get_cell_types() - # Create weights DataFrame - weights_df = pd.DataFrame( - weights, - index=sample_indices, - columns=cell_types - ) - # Assign cell types sample_data.obs['cell_type'] = weights_df.idxmax(axis=1) @@ -363,7 +357,7 @@ def main(): weights_plot_path = figures_dir / f"{args.output}{sample_suffix}_weights.png" fig = plot_spatial_weights( coords=sample_data.obsm['spatial'], - weights=weights, + weights=weights_df, cell_types=cell_types, save_path=str(weights_plot_path), ) diff --git a/tests/_helpers.py b/tests/_helpers.py new file mode 100644 index 0000000..ef4a6ba --- /dev/null +++ b/tests/_helpers.py @@ -0,0 +1,68 @@ +"""Shared synthetic-data builders reused across test_preprocessing.py and +test_core.py. Not a test module itself -- no tests live here. +""" + +import numpy as np +import pandas as pd +from anndata import AnnData + + +def _make_expression(n_cells, n_genes, seed=0): + rng = np.random.default_rng(seed) + return rng.integers(1, 10, size=(n_cells, n_genes)).astype(np.float64) + + +def _make_obs(subtype_sizes): + """ + subtype_sizes : dict mapping major type name to a list of subtype cell + counts, e.g. {"A": [40, 40, 20]} makes major type "A" with subtypes + "A_0", "A_1", "A_2" of the given sizes. + """ + majors = [] + subtypes = [] + for major, sizes in subtype_sizes.items(): + for i, n in enumerate(sizes): + majors.extend([major] * n) + subtypes.extend([f"{major}_{i}"] * n) + return pd.DataFrame({"major": majors, "subtype": subtypes}) + + +def _make_ref_adata(obs, n_genes=20, seed=0): + """Wrap a given obs DataFrame (built via _make_obs or inline) into an AnnData.""" + obs = obs.reset_index(drop=True) + obs.index = obs.index.astype(str) + X = _make_expression(len(obs), n_genes, seed=seed) + var = pd.DataFrame(index=[f"gene{i}" for i in range(n_genes)]) + return AnnData(X=X, obs=obs, var=var) + + +def _make_spatial_adata(n_spots, n_genes, seed=0): + X = _make_expression(n_spots, n_genes, seed=seed) + adata = AnnData(X=X, obs=pd.DataFrame(index=[str(i) for i in range(n_spots)])) + adata.var_names = [f"gene{i}" for i in range(n_genes)] + adata.obsm["spatial"] = np.random.default_rng(seed).uniform(0, 10, size=(n_spots, 2)) + return adata + + +def _make_clustered_spatial_adata(n_clusters, spots_per_cluster, n_genes=5, cluster_spacing=1000.0, jitter=0.3, seed=0): + """ + Spots in tight, well-separated clusters. Every spot in a cluster shares + an identical expression profile (cosine similarity 1.0 to its + cluster-mates), so which pairs form is governed purely by spatial + radius, not by expression-threshold luck. + """ + rng = np.random.default_rng(seed) + coords_parts = [] + X_parts = [] + for c in range(n_clusters): + anchor = np.array([c * cluster_spacing, 0.0]) + coords_parts.append(anchor + rng.uniform(-jitter, jitter, size=(spots_per_cluster, 2))) + base = _make_expression(1, n_genes, seed=seed * 100 + c) + X_parts.append(np.repeat(base, spots_per_cluster, axis=0)) + coords = np.vstack(coords_parts) + X = np.vstack(X_parts) + n_spots = X.shape[0] + adata = AnnData(X=X, obs=pd.DataFrame(index=[str(i) for i in range(n_spots)])) + adata.var_names = [f"gene{i}" for i in range(n_genes)] + adata.obsm["spatial"] = coords + return adata diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..290cc21 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,3 @@ +import matplotlib + +matplotlib.use("Agg") diff --git a/tests/test_core.py b/tests/test_core.py new file mode 100644 index 0000000..bd0cf90 --- /dev/null +++ b/tests/test_core.py @@ -0,0 +1,1041 @@ +import numpy as np +import pandas as pd +import pytest +import torch +from anndata import AnnData +from scipy.sparse import csr_matrix, issparse, triu + +from dotpy.core import ( + DOT, + MODE_PRESETS, + _clusters_at_level, + _default_weights_to_lambdas, + _ref_to_anndata, + _resolve_lambdas, + _safe_log2, + _spatial_to_anndata, + _sqrt_env, + _sqrt_env_grad, + _validate_weights, +) +from dotpy.preprocessing import setup_reference, setup_spatial +from _helpers import ( + _make_clustered_spatial_adata, + _make_expression, + _make_ref_adata, + _make_spatial_adata, +) + + +def _ref_dict_from_adata(ref_adata): + """Convert a real setup_reference() AnnData output into the legacy dict shape DOT still accepts.""" + coarsest_key = ref_adata.uns["level_keys"][0] + labels = ref_adata.obs[coarsest_key].values + clusters = {ct: np.where(labels == ct)[0] for ct in pd.unique(labels)} + return { + "X_sparse": ref_adata.X, + "clusters": clusters, + "ratios": dict(ref_adata.uns["ratios"]), + "genes": ref_adata.var_names.values, + } + + +def _spatial_dict_from_adata(spatial_adata): + """Convert a real setup_spatial() AnnData output into the legacy dict shape DOT still accepts.""" + d = { + "X_sparse": spatial_adata.X, + "coords": spatial_adata.obsm["spatial"], + "genes": spatial_adata.var_names.values, + } + if "dot_spatial_pairs" in spatial_adata.obsp: + mat = triu(spatial_adata.obsp["dot_spatial_pairs"], k=1).tocoo() + d["pairs"] = {"i": mat.row, "j": mat.col, "w": mat.data} + return d + + +def _make_dot_inputs(n_cells=60, n_spots=20, n_genes=10, seed=0, with_pairs=False): + """Build a small, real ref/spatial AnnData pair via setup_reference()/setup_spatial(), suitable for DOT().""" + obs = pd.DataFrame({"cell_type": ["A"] * (n_cells // 2) + ["B"] * (n_cells // 2)}) + ref_input = _make_ref_adata(obs, n_genes=n_genes, seed=seed) + ref_adata = setup_reference( + ref_input, cell_type_key="cell_type", subcluster_size=3, + random_state=0, max_input_genes=100, + ) + + if with_pairs: + spots_per_cluster = max(2, n_spots // 4) + spatial_input = _make_clustered_spatial_adata( + n_clusters=4, spots_per_cluster=spots_per_cluster, n_genes=n_genes, seed=seed, + ) + spatial_adata = setup_spatial( + spatial_input, radius=2.0, th_gene_low=0.0, th_gene_high=1.0, + ) + else: + spatial_input = _make_spatial_adata(n_spots, n_genes, seed=seed) + spatial_adata = setup_spatial( + spatial_input, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0, + ) + + return ref_adata, spatial_adata + + +class TestRefToAnnData: + def test_anndata_passthrough_is_identity(self): + ref_adata, _ = _make_dot_inputs() + result = _ref_to_anndata(ref_adata) + assert result is ref_adata + + def test_dict_input_converts_correctly(self): + ref_adata, _ = _make_dot_inputs() + ref_dict = _ref_dict_from_adata(ref_adata) + with pytest.warns(FutureWarning, match="deprecated"): + result = _ref_to_anndata(ref_dict) + + assert isinstance(result, AnnData) + assert result.uns["level_keys"] == ["major"] + assert np.allclose(np.asarray(result.X), np.asarray(ref_adata.X)) + assert list(result.var_names) == list(ref_adata.var_names) + assert dict(result.uns["ratios"]) == dict(ref_adata.uns["ratios"]) + + expected_major = ref_adata.obs[ref_adata.uns["level_keys"][0]].values + assert list(result.obs["major"]) == list(expected_major) + + def test_dict_input_raises_exactly_one_warning(self, recwarn): + ref_adata, _ = _make_dot_inputs() + ref_dict = _ref_dict_from_adata(ref_adata) + _ref_to_anndata(ref_dict) + future_warnings = [w for w in recwarn if issubclass(w.category, FutureWarning)] + assert len(future_warnings) == 1 + + def test_invalid_type_raises(self): + with pytest.raises(TypeError, match="ref must be"): + _ref_to_anndata([1, 2, 3]) + + +class TestSpatialToAnnData: + def test_anndata_passthrough_is_identity(self): + _, spatial_adata = _make_dot_inputs() + result = _spatial_to_anndata(spatial_adata) + assert result is spatial_adata + + def test_dict_input_converts_correctly_without_pairs(self): + _, spatial_adata = _make_dot_inputs(with_pairs=False) + spatial_dict = _spatial_dict_from_adata(spatial_adata) + with pytest.warns(FutureWarning, match="deprecated"): + result = _spatial_to_anndata(spatial_dict) + + assert isinstance(result, AnnData) + assert np.allclose(np.asarray(result.X), np.asarray(spatial_adata.X)) + assert list(result.var_names) == list(spatial_adata.var_names) + assert np.allclose(result.obsm["spatial"], spatial_adata.obsm["spatial"]) + assert "dot_spatial_pairs" not in result.obsp + + def test_dict_input_converts_correctly_with_pairs(self): + _, spatial_adata = _make_dot_inputs(with_pairs=True) + assert "dot_spatial_pairs" in spatial_adata.obsp # sanity check on the fixture itself + + spatial_dict = _spatial_dict_from_adata(spatial_adata) + with pytest.warns(FutureWarning, match="deprecated"): + result = _spatial_to_anndata(spatial_dict) + + assert "dot_spatial_pairs" in result.obsp + assert np.allclose( + result.obsp["dot_spatial_pairs"].toarray(), + spatial_adata.obsp["dot_spatial_pairs"].toarray(), + ) + + def test_dict_input_raises_exactly_one_warning(self, recwarn): + _, spatial_adata = _make_dot_inputs() + spatial_dict = _spatial_dict_from_adata(spatial_adata) + _spatial_to_anndata(spatial_dict) + future_warnings = [w for w in recwarn if issubclass(w.category, FutureWarning)] + assert len(future_warnings) == 1 + + def test_invalid_type_raises(self): + with pytest.raises(TypeError, match="spatial must be"): + _spatial_to_anndata(42) + + +class TestClustersAtLevel: + def test_defaults_to_coarsest_level(self): + obs = pd.DataFrame({ + "major": ["A", "A", "B", "B", "B"], + "subtype": ["A_0", "A_1", "B_0", "B_0", "B_1"], + }) + obs.index = obs.index.astype(str) + adata = AnnData(X=_make_expression(5, 4, seed=0), obs=obs) + adata.uns["level_keys"] = ["major", "subtype"] + + result = _clusters_at_level(adata) + assert set(result.keys()) == {"A", "B"} + assert sorted(result["A"].tolist()) == [0, 1] + assert sorted(result["B"].tolist()) == [2, 3, 4] + + def test_groups_by_explicit_finer_level(self): + obs = pd.DataFrame({ + "major": ["A", "A", "B", "B", "B"], + "subtype": ["A_0", "A_1", "B_0", "B_0", "B_1"], + }) + obs.index = obs.index.astype(str) + adata = AnnData(X=_make_expression(5, 4, seed=0), obs=obs) + adata.uns["level_keys"] = ["major", "subtype"] + + result = _clusters_at_level(adata, level="subtype") + assert set(result.keys()) == {"A_0", "A_1", "B_0", "B_1"} + assert sorted(result["B_0"].tolist()) == [2, 3] + + def test_single_level_hierarchy(self): + obs = pd.DataFrame({"major": ["A", "B", "A"]}) + obs.index = obs.index.astype(str) + adata = AnnData(X=_make_expression(3, 4, seed=1), obs=obs) + adata.uns["level_keys"] = ["major"] + + result = _clusters_at_level(adata) + assert sorted(result["A"].tolist()) == [0, 2] + assert sorted(result["B"].tolist()) == [1] + + def test_invalid_level_raises(self): + obs = pd.DataFrame({"major": ["A", "B", "A"]}) + obs.index = obs.index.astype(str) + adata = AnnData(X=_make_expression(3, 4, seed=2), obs=obs) + adata.uns["level_keys"] = ["major"] + + with pytest.raises(ValueError, match="level must be one of"): + _clusters_at_level(adata, level="bogus") + + +class TestSafeLog2: + """Pins _safe_log2 to R's safelog2 boundary values.""" + + def test_zero_maps_to_minus_20(self): + out = _safe_log2(torch.tensor([0.0])) + assert out.item() == pytest.approx(-20.0) + + def test_negative_maps_to_zero(self): + out = _safe_log2(torch.tensor([-1.0, -100.0])) + assert out.tolist() == pytest.approx([0.0, 0.0]) + + def test_positive_inf_maps_to_minus_20(self): + out = _safe_log2(torch.tensor([float("inf")])) + assert out.item() == pytest.approx(-20.0) + + def test_normal_positive_values_match_log2_exactly(self): + out = _safe_log2(torch.tensor([1.0, 2.0, 4.0, 0.5])) + assert out.tolist() == pytest.approx([0.0, 1.0, 2.0, -1.0]) + + def test_matches_r_safelog2_on_a_mixed_vector(self): + # Same vector run through R's safelog2 by hand: log2(0)->-20, + # log2(-1)->0 (NaN->0), log2(4)=2, log2(1)=0. + out = _safe_log2(torch.tensor([0.0, -1.0, 4.0, 1.0])) + assert out.tolist() == pytest.approx([-20.0, 0.0, 2.0, 0.0]) + + +class TestSqrtEnv: + """Pins _sqrt_env/_sqrt_env_grad's piecewise linear/sqrt envelope.""" + + _THRESHOLD = 4e-4 + + def test_below_threshold_matches_linear_formula(self): + v = torch.tensor([0.0, 1e-4]) + out = _sqrt_env(v) + assert out.tolist() == pytest.approx([0.01, 0.0125]) + + def test_at_or_above_threshold_matches_sqrt(self): + v = torch.tensor([self._THRESHOLD, 1.0, 4.0]) + out = _sqrt_env(v) + assert out.tolist() == pytest.approx([0.02, 1.0, 2.0]) + + def test_continuous_at_threshold_boundary(self): + just_below = _sqrt_env(torch.tensor([self._THRESHOLD - 1e-8])) + at_threshold = _sqrt_env(torch.tensor([self._THRESHOLD])) + assert just_below.item() == pytest.approx(at_threshold.item(), abs=1e-5) + + def test_grad_below_threshold_is_constant_slope(self): + v = torch.tensor([0.0, 1e-4]) + out = _sqrt_env_grad(v) + assert out.tolist() == pytest.approx([25.0, 25.0]) + + def test_grad_at_or_above_threshold_matches_derivative_formula(self): + v = torch.tensor([1.0, 4.0]) + out = _sqrt_env_grad(v) + assert out.tolist() == pytest.approx([0.5, 0.25], abs=1e-5) + + def test_grad_at_threshold_is_25(self): + # The envelope is constructed so the linear piece is exactly + # tangent to sqrt at the boundary: 0.5/sqrt(4e-4) = 25.0, matching + # the linear slope. Checked directly, not via finite differences, + # since a finite-difference window straddling this exact point + # picks up sqrt's curvature rather than the pointwise derivative. + out = _sqrt_env_grad(torch.tensor([self._THRESHOLD])) + assert out.item() == pytest.approx(25.0) + + def test_grad_matches_finite_difference_of_sqrt_env(self): + eps = 1e-4 + for v0 in [0.0, 1e-4, 0.01, 1.0, 4.0]: + v = torch.tensor([v0]) + numerical = (_sqrt_env(v + eps) - _sqrt_env(v - eps)) / (2 * eps) + analytical = _sqrt_env_grad(v) + assert numerical.item() == pytest.approx(analytical.item(), abs=1e-2) + + +class TestValidateWeights: + def test_all_valid_does_not_raise(self): + _validate_weights(1.0, 1.0, 0.01, 0.0, 0.6, 0.0) # no raise + + def test_negative_weight_raises(self): + with pytest.raises(ValueError, match="gene_weight must be >= 0"): + _validate_weights(-1.0, 1.0, 0.01, 0.0, 0.6, 0.0) + + def test_all_zero_raises(self): + with pytest.raises(ValueError, match="must be > 0"): + _validate_weights(0.0, 0.0, 0.0, 0.0, 0.6, 0.0) + + def test_sparsity_coef_above_one_raises(self): + with pytest.raises(ValueError, match="sparsity_coef must be in"): + _validate_weights(1.0, 1.0, 0.01, 0.0, 1.5, 0.0) + + def test_sparsity_coef_below_zero_raises(self): + with pytest.raises(ValueError, match="sparsity_coef must be in"): + _validate_weights(1.0, 1.0, 0.01, 0.0, -0.1, 0.0) + + def test_sparsity_coef_boundary_values_do_not_raise(self): + _validate_weights(1.0, 1.0, 0.01, 0.0, 0.0, 0.0) # no raise + _validate_weights(1.0, 1.0, 0.01, 0.0, 1.0, 0.0) # no raise + + +class TestDefaultWeightsToLambdas: + def test_matches_hand_computed_values(self): + weights = { + "gene_weight": 2.0, "spot_weight": 4.0, "spatial_weight": 3.0, + "ratios_weight": 6.0, "sparsity_coef": 0.25, "cluster_weight": 5.0, + } + result = _default_weights_to_lambdas( + weights, S=10, G=5, C=4, max_size=2, n_pairs=4, has_pairs=True, + ) + assert result["l_a"] == pytest.approx(3.0) + assert result["l_g"] == pytest.approx(4.0) + assert result["l_i"] == pytest.approx(3.0) + assert result["l_sp"] == pytest.approx(0.5) + assert result["l_s"] == pytest.approx(3.75) + assert result["l_c"] == pytest.approx(12.5) + + def test_has_pairs_false_gives_zero_l_s(self): + weights = { + "gene_weight": 1.0, "spot_weight": 1.0, "spatial_weight": 5.0, + "ratios_weight": 0.0, "sparsity_coef": 0.5, "cluster_weight": 0.0, + } + result = _default_weights_to_lambdas( + weights, S=10, G=5, C=4, max_size=1, n_pairs=0, has_pairs=False, + ) + assert result["l_s"] == 0.0 + + def test_sparsity_coef_extremes_fold_into_l_i(self): + weights = { + "gene_weight": 1.0, "spot_weight": 2.0, "spatial_weight": 0.0, + "ratios_weight": 0.0, "sparsity_coef": 1.0, "cluster_weight": 0.0, + } + result = _default_weights_to_lambdas( + weights, S=10, G=5, C=4, max_size=1, n_pairs=0, has_pairs=False, + ) + assert result["l_i"] == pytest.approx(0.0) + + weights["sparsity_coef"] = 0.0 + result = _default_weights_to_lambdas( + weights, S=10, G=5, C=4, max_size=1, n_pairs=0, has_pairs=False, + ) + assert result["l_i"] == pytest.approx(2.0) + + +class TestResolveLambdas: + _base_weights = { + "gene_weight": 1.0, "spot_weight": 1.0, "spatial_weight": 0.01, + "ratios_weight": 0.0, "sparsity_coef": 0.6, "cluster_weight": 0.0, + } + + def test_no_override_returns_computed_lambdas(self): + result = _resolve_lambdas( + self._base_weights, None, _default_weights_to_lambdas, + S=10, G=5, C=4, max_size=1, n_pairs=4, has_pairs=True, + ) + expected = _default_weights_to_lambdas( + self._base_weights, S=10, G=5, C=4, max_size=1, n_pairs=4, has_pairs=True, + ) + assert result == expected + + def test_override_applies(self): + result = _resolve_lambdas( + self._base_weights, {"l_g": 99.0}, _default_weights_to_lambdas, + S=10, G=5, C=4, max_size=1, n_pairs=4, has_pairs=True, + ) + assert result["l_g"] == 99.0 + + def test_unknown_lambda_key_raises(self): + with pytest.raises(ValueError, match="unknown keys"): + _resolve_lambdas( + self._base_weights, {"l_x": 1.0}, _default_weights_to_lambdas, + S=10, G=5, C=4, max_size=1, n_pairs=4, has_pairs=True, + ) + + def test_custom_callable_missing_key_raises(self): + def incomplete(weights, *, S, G, C, max_size, n_pairs, has_pairs): + return {"l_a": 0.0, "l_g": 0.0, "l_i": 0.0, "l_sp": 0.0, "l_c": 0.0} # missing l_s + + with pytest.raises(ValueError, match="must return all of"): + _resolve_lambdas( + self._base_weights, None, incomplete, + S=10, G=5, C=4, max_size=1, n_pairs=4, has_pairs=True, + ) + + def test_default_transform_no_pairs_does_not_warn(self, recwarn): + result = _resolve_lambdas( + self._base_weights, None, _default_weights_to_lambdas, + S=10, G=5, C=4, max_size=1, n_pairs=0, has_pairs=False, + ) + assert result["l_s"] == 0.0 + user_warnings = [w for w in recwarn if issubclass(w.category, UserWarning)] + assert len(user_warnings) == 0 + + def test_explicit_override_no_pairs_warns_and_zeroes(self): + with pytest.warns(UserWarning, match="no spatial neighbour pairs"): + result = _resolve_lambdas( + self._base_weights, {"l_s": 5.0}, _default_weights_to_lambdas, + S=10, G=5, C=4, max_size=1, n_pairs=0, has_pairs=False, + ) + assert result["l_s"] == 0.0 + + def test_noncompliant_custom_callable_no_pairs_warns_and_zeroes(self): + def bad_transform(weights, *, S, G, C, max_size, n_pairs, has_pairs): + return {"l_a": 0.0, "l_g": 0.0, "l_i": 0.0, "l_sp": 0.0, "l_c": 0.0, "l_s": 3.0} + + with pytest.warns(UserWarning, match="no spatial neighbour pairs"): + result = _resolve_lambdas( + self._base_weights, None, bad_transform, + S=10, G=5, C=4, max_size=1, n_pairs=0, has_pairs=False, + ) + assert result["l_s"] == 0.0 + + def test_unknown_lambda_key_includes_l_c_in_valid_set(self): + # Sanity check that l_c is a recognized override key now, not + # rejected as unknown. + result = _resolve_lambdas( + self._base_weights, {"l_c": 7.0}, _default_weights_to_lambdas, + S=10, G=5, C=4, max_size=1, n_pairs=4, has_pairs=True, + ) + assert result["l_c"] == 7.0 + + +class TestDOTOptConfig: + def test_highres_default_weights_and_lambdas(self): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + w = dot.opt_config["weights"] + assert w == { + "gene_weight": 1.0, "spot_weight": 1.0, "spatial_weight": 0.01, + "ratios_weight": 0.0, "sparsity_coef": 0.6, "cluster_weight": 0.0, + } + assert dot.opt_config["max_size"] == 1 + assert dot.opt_config["lambda_overrides"] == set() + assert dot.opt_config["lambdas"]["l_c"] == 0.0 + + def test_explicit_cluster_weight_overrides_preset(self): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + dot = DOT(spatial=spatial_adata, ref=ref_adata, cluster_weight=2.0, ls_solution=False) + assert dot.opt_config["weights"]["cluster_weight"] == 2.0 + + def test_lowres_default_weights_and_lambdas(self): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + dot = DOT(spatial=spatial_adata, ref=ref_adata, mode="lowres", ls_solution=False) + w = dot.opt_config["weights"] + assert w["spot_weight"] == 0.25 + assert w["sparsity_coef"] == 0.4 + assert dot.opt_config["max_size"] == 20 + + def test_spot_weight_default_depends_on_resolved_max_size(self): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + dot = DOT(spatial=spatial_adata, ref=ref_adata, mode="lowres", max_size=1, ls_solution=False) + assert dot.opt_config["weights"]["spot_weight"] == 1.0 + + def test_explicit_weight_overrides_preset(self): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + dot = DOT(spatial=spatial_adata, ref=ref_adata, spatial_weight=0.5, ls_solution=False) + assert dot.opt_config["weights"]["spatial_weight"] == 0.5 + + def test_invalid_mode_raises(self): + ref_adata, spatial_adata = _make_dot_inputs() + with pytest.raises(ValueError, match="mode must be one of"): + DOT(spatial=spatial_adata, ref=ref_adata, mode="bogus", ls_solution=False) + + def test_lambdas_override_tracked(self): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + dot = DOT(spatial=spatial_adata, ref=ref_adata, lambdas={"l_g": 2.0}, ls_solution=False) + assert dot.opt_config["lambdas"]["l_g"] == 2.0 + assert dot.opt_config["lambda_overrides"] == {"l_g"} + + def test_custom_weights_to_lambdas_used(self): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + + def custom(weights, *, S, G, C, max_size, n_pairs, has_pairs): + return {"l_a": 0.0, "l_g": 0.0, "l_i": 0.0, "l_sp": 0.0, "l_c": 0.0, "l_s": 0.0} + + dot = DOT(spatial=spatial_adata, ref=ref_adata, weights_to_lambdas=custom, ls_solution=False) + assert dot.opt_config["lambdas"] == { + "l_a": 0.0, "l_g": 0.0, "l_i": 0.0, "l_sp": 0.0, "l_c": 0.0, "l_s": 0.0, + } + + def test_custom_weights_to_lambdas_missing_key_raises(self): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + + def incomplete(weights, *, S, G, C, max_size, n_pairs, has_pairs): + return {"l_a": 0.0, "l_g": 0.0, "l_i": 0.0, "l_sp": 0.0} + + with pytest.raises(ValueError, match="must return all of"): + DOT(spatial=spatial_adata, ref=ref_adata, weights_to_lambdas=incomplete, ls_solution=False) + + +class TestPrintConfig: + def test_runs_without_error_and_prints_something(self, capsys): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + dot.print_config() + out = capsys.readouterr().out + assert out != "" + + def test_default_transform_shows_weight_param_column(self, capsys): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + dot.print_config() + out = capsys.readouterr().out + assert "weight param" in out + assert "spot_weight * (1 - sparsity_coef)" in out + assert "spot_weight * sparsity_coef" in out + assert "cluster_weight" in out + assert "l_c" in out + + def test_custom_transform_omits_weight_param_and_shows_disclaimer(self, capsys): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + + def my_custom_transform(weights, *, S, G, C, max_size, n_pairs, has_pairs): + return {"l_a": 0.1, "l_g": 1.5, "l_i": 0.3, "l_sp": 0.05, "l_c": 0.0, "l_s": 0.0} + + dot = DOT( + spatial=spatial_adata, ref=ref_adata, + weights_to_lambdas=my_custom_transform, ls_solution=False, + ) + dot.print_config() + out = capsys.readouterr().out + assert "weight param" not in out + assert "my_custom_transform" in out + assert "not the default formula" in out + + def test_override_shows_user_set_note(self, capsys): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + dot = DOT(spatial=spatial_adata, ref=ref_adata, lambdas={"l_g": 2.0}, ls_solution=False) + dot.print_config() + out = capsys.readouterr().out + assert "user-set directly" in out + assert "2.0000" in out + + def test_inactive_note_only_appears_for_l_s_without_pairs(self, capsys): + ref_adata1, spatial_adata_no_pairs = _make_dot_inputs(with_pairs=False) + dot_no_pairs = DOT(spatial=spatial_adata_no_pairs, ref=ref_adata1, ls_solution=False) + dot_no_pairs.print_config() + out_no_pairs = capsys.readouterr().out + assert "inactive (no spatial pairs)" in out_no_pairs + + ref_adata2, spatial_adata_pairs = _make_dot_inputs(with_pairs=True) + dot_pairs = DOT(spatial=spatial_adata_pairs, ref=ref_adata2, ls_solution=False) + dot_pairs.print_config() + out_pairs = capsys.readouterr().out + assert "inactive" not in out_pairs + + +class TestDOTInit: + def test_anndata_inputs_produce_anndata_attributes(self, recwarn): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + + assert isinstance(dot.ref, AnnData) + assert isinstance(dot.spatial, AnnData) + future_warnings = [w for w in recwarn if issubclass(w.category, FutureWarning)] + assert len(future_warnings) == 0 + + def test_dict_inputs_raise_two_future_warnings(self, recwarn): + ref_adata, spatial_adata = _make_dot_inputs() + ref_dict = _ref_dict_from_adata(ref_adata) + spatial_dict = _spatial_dict_from_adata(spatial_adata) + DOT(spatial=spatial_dict, ref=ref_dict, ls_solution=False) + + future_warnings = [w for w in recwarn if issubclass(w.category, FutureWarning)] + assert len(future_warnings) == 2 + + def test_mixed_inputs_raise_one_future_warning(self, recwarn): + ref_adata, spatial_adata = _make_dot_inputs() + ref_dict = _ref_dict_from_adata(ref_adata) + DOT(spatial=spatial_adata, ref=ref_dict, ls_solution=False) + + future_warnings = [w for w in recwarn if issubclass(w.category, FutureWarning)] + assert len(future_warnings) == 1 + + def test_no_common_genes_raises(self): + ref_adata, spatial_adata = _make_dot_inputs() + spatial_adata = spatial_adata.copy() + spatial_adata.var_names = [f"other_gene{i}" for i in range(spatial_adata.n_vars)] + with pytest.raises(ValueError, match="No common genes"): + DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + + def test_invalid_ref_type_raises(self): + _, spatial_adata = _make_dot_inputs() + with pytest.raises(TypeError, match="ref must be"): + DOT(spatial=spatial_adata, ref="not valid", ls_solution=False) + + def test_device_defaults_to_cpu(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + assert dot.device == "cpu" + + def test_ls_solution_true_populates_solution(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=True) + assert dot.solution is not None + assert dot.solution.shape == (dot.ref.n_obs, dot.spatial.n_obs) + + def test_ls_solution_false_leaves_solution_none(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + assert dot.solution is None + + def test_ls_solution_matches_between_sparse_and_dense_input(self): + ref_adata, spatial_adata = _make_dot_inputs() + + dot_dense = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + sol_dense = dot_dense._ls_solution() + + dot_sparse = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + dot_sparse.ref.X = csr_matrix(dot_sparse.ref.X) + dot_sparse.spatial.X = csr_matrix(dot_sparse.spatial.X) + sol_sparse = dot_sparse._ls_solution() + + assert np.allclose(sol_dense, sol_sparse) + + def test_get_cell_types_returns_major_labels(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + assert set(dot.get_cell_types()) == {"A", "B"} + + def test_get_cell_types_with_explicit_level_returns_finer_labels(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + _, finest_key = ref_adata.uns["level_keys"] + fine_labels = dot.get_cell_types(level=finest_key) + assert set(fine_labels) == set(ref_adata.obs[finest_key]) + + def test_get_cell_types_invalid_level_raises(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + with pytest.raises(ValueError, match="level must be one of"): + dot.get_cell_types(level="bogus") + + def test_get_cell_types_returns_list_with_no_duplicates(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + cell_types = dot.get_cell_types() + assert isinstance(cell_types, list) + assert len(cell_types) == len(set(cell_types)) + + def test_gene_alignment_keeps_only_common_genes_in_order(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + common = np.intersect1d(np.asarray(ref_adata.var_names), np.asarray(spatial_adata.var_names)) + assert list(dot.ref.var_names) == list(common) + assert list(dot.spatial.var_names) == list(common) + + +class TestDOTFitOutputFormat: + """ + Checks that fit() runs and produces correctly-shaped, correctly + normalized output for both input paths. Does not check numerical + correctness of the deconvolution result itself (convergence quality, + loss-term correctness) + """ + + def test_fit_runs_and_produces_correctly_shaped_weights(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot.fit(iterations=5, verbose=False) + + weights = dot.get_weights() + assert weights.shape == (spatial_adata.n_obs, len(dot.get_cell_types())) + assert list(weights.index) == list(dot.spatial.obs_names) + assert list(weights.columns) == dot.get_cell_types() + + def test_fit_without_ls_solution_runs(self): + # Exercises the cold-start Yt initialization path in + # _run_optimisation (the "if Yt is None" branch), only reached + # when ls_solution=False leaves self.solution as None going into + # fit() -- never exercised by any other test, which all use the + # default ls_solution=True. + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False, device="cpu") + assert dot.solution is None + dot.fit(iterations=5, verbose=False) + + weights = dot.get_weights() + assert weights.shape == (spatial_adata.n_obs, len(dot.get_cell_types())) + assert np.isfinite(weights.values).all() + + def test_fit_rescues_near_zero_ls_solution_column(self): + # Exercises the near-zero-column rescue in the warm-start branch + # (Yt_cand[:, small] = 1.0 / C), for a spot whose LS-solution + # column sums to < 1e-3. Hand-crafted directly on dot.solution + # rather than relying on the real LS solve to happen to produce + # one, for a deterministic trigger. + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False, device="cpu") + C, S = dot.ref.n_obs, dot.spatial.n_obs + solution = np.full((C, S), 0.5, dtype=np.float32) + solution[:, 0] = 1e-6 # column sums to ~C*1e-6, well under 1e-3 + dot.solution = solution + dot.fit(iterations=5, verbose=False) + assert np.isfinite(dot.get_weights().values).all() + + def test_fit_rescales_ls_solution_with_low_sparsity_coef(self): + # Exercises the sparsity_coef <= 0.5 rescaling branch (max_size/ + # min_size-based, vs. the >0.5 branch's simple normalize-to-1), + # only reached with lowres mode's default sparsity_coef (0.4). + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, mode="lowres", ls_solution=False, device="cpu") + C, S = dot.ref.n_obs, dot.spatial.n_obs + solution = np.full((C, S), 0.5, dtype=np.float32) + solution[:, 0] = 100.0 # column sum far exceeds max_size=20 + dot.solution = solution + dot.fit(iterations=5, verbose=False) + assert np.isfinite(dot.get_weights().values).all() + + def test_fit_mixes_uniform_weight_into_exact_zero_entries(self): + # Exercises the (Yt_cand == 0).any() fallback -- a literal zero + # surviving clamp+rescale gets blended toward uniform, avoiding a + # degenerate all-zero entry feeding into the optimisation. + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False, device="cpu") + C, S = dot.ref.n_obs, dot.spatial.n_obs + solution = np.full((C, S), 0.5, dtype=np.float32) + solution[0, 0] = 0.0 # exact zero, in a column that isn't "small" overall + dot.solution = solution + dot.fit(iterations=5, verbose=False) + assert np.isfinite(dot.get_weights().values).all() + + def test_fit_rejects_old_removed_params(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + with pytest.raises(TypeError): + dot.fit(mode="highres") + + def test_lambda_override_actually_affects_fit(self): + ref_adata, spatial_adata = _make_dot_inputs() + + dot_default = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot_default.fit(iterations=5, verbose=False) + + dot_override = DOT(spatial=spatial_adata, ref=ref_adata, lambdas={"l_g": 0.0}, device="cpu") + dot_override.fit(iterations=5, verbose=False) + + assert not np.allclose(dot_default.get_weights().values, dot_override.get_weights().values) + + def test_ratios_weight_actually_affects_fit(self): + # ratios_weight/l_a defaults to 0.0 in both mode presets and was + # never previously set to a nonzero value in any test, so the + # abundance-matching term had never actually been exercised. + ref_adata, spatial_adata = _make_dot_inputs() + + dot_default = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot_default.fit(iterations=5, verbose=False) + + dot_ratios = DOT(spatial=spatial_adata, ref=ref_adata, ratios_weight=1.0, device="cpu") + dot_ratios.fit(iterations=5, verbose=False) + + assert not np.allclose(dot_default.get_weights().values, dot_ratios.get_weights().values) + + def test_cluster_weight_actually_affects_fit(self): + ref_adata, spatial_adata = _make_dot_inputs() + + dot_default = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot_default.fit(iterations=5, verbose=False) + + dot_cw = DOT(spatial=spatial_adata, ref=ref_adata, cluster_weight=10.0, device="cpu") + dot_cw.fit(iterations=5, verbose=False) + + assert not np.allclose(dot_default.get_weights().values, dot_cw.get_weights().values) + + def test_get_weights_at_finer_level_has_more_columns(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot.fit(iterations=5, verbose=False) + + coarsest_key, finest_key = ref_adata.uns["level_keys"] + coarse_weights = dot.get_weights(level=coarsest_key) + fine_weights = dot.get_weights(level=finest_key) + + assert list(coarse_weights.columns) == dot.get_cell_types(level=coarsest_key) + assert list(fine_weights.columns) == dot.get_cell_types(level=finest_key) + assert fine_weights.shape[1] >= coarse_weights.shape[1] + + def test_finer_level_weights_sum_to_coarser_level_weights(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot.fit(iterations=5, verbose=False) + + coarsest_key, finest_key = ref_adata.uns["level_keys"] + coarse_weights = dot.get_weights(level=coarsest_key, normalize=False) + fine_weights = dot.get_weights(level=finest_key, normalize=False) + parent_of = ref_adata.obs.set_index(finest_key)[coarsest_key] + + for ct in coarse_weights.columns: + children = [c for c in fine_weights.columns if parent_of[c] == ct] + summed = fine_weights[children].sum(axis=1) + assert np.allclose(summed, coarse_weights[ct]) + + def test_get_weights_default_normalizes_rows_to_one(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot.fit(iterations=5, verbose=False) + + weights = dot.get_weights() + row_sums = weights.sum(axis=1) + assert np.allclose(row_sums, 1.0, atol=1e-3) + + def test_get_weights_normalize_false_returns_raw_weights(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot.fit(iterations=5, verbose=False) + + raw = dot.get_weights(normalize=False) + assert np.allclose(raw.values, dot.weights) + + def test_get_weights_handles_all_zero_row_without_nan(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot.fit(iterations=5, verbose=False) + + dot.weights[0, :] = 0.0 # force one spot to have zero total weight + weights = dot.get_weights(normalize=True) + assert not np.isnan(weights.values).any() + assert np.allclose(weights.iloc[0].values, 0.0) + + def test_get_weights_invalid_level_raises(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot.fit(iterations=5, verbose=False) + with pytest.raises(ValueError, match="level must be one of"): + dot.get_weights(level="bogus") + + def test_get_weights_explicit_coarsest_level_matches_default(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot.fit(iterations=5, verbose=False) + + coarsest_key = ref_adata.uns["level_keys"][0] + w_default = dot.get_weights() + w_explicit = dot.get_weights(level=coarsest_key) + assert list(w_default.columns) == list(w_explicit.columns) + assert np.allclose(w_default.values, w_explicit.values) + + def test_get_weights_does_not_mutate_self_weights_or_solution(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot.fit(iterations=5, verbose=False) + + weights_before = dot.weights.copy() + solution_before = dot.solution.copy() + finest_key = ref_adata.uns["level_keys"][1] + + dot.get_weights(normalize=True) + dot.get_weights(normalize=False) + dot.get_weights(level=finest_key, normalize=True) + + assert np.allclose(dot.weights, weights_before) + assert np.allclose(dot.solution, solution_before) + + def test_dict_input_dot_only_supports_major_level(self): + ref_adata, spatial_adata = _make_dot_inputs() + ref_dict = _ref_dict_from_adata(ref_adata) + spatial_dict = _spatial_dict_from_adata(spatial_adata) + + with pytest.warns(FutureWarning, match="deprecated"): + dot = DOT(spatial=spatial_dict, ref=ref_dict, device="cpu") + dot.fit(iterations=5, verbose=False) + + assert dot.get_cell_types() == dot.get_cell_types(level="major") + w_default = dot.get_weights() + w_major = dot.get_weights(level="major") + assert list(w_default.columns) == list(w_major.columns) + assert np.allclose(w_default.values, w_major.values) + + with pytest.raises(ValueError, match="level must be one of"): + dot.get_weights(level="subtype") + + def test_get_weights_before_fit_raises(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False) + with pytest.raises(ValueError, match="Not fitted"): + dot.get_weights() + + def test_fit_with_spatial_pairs_runs(self): + ref_adata, spatial_adata = _make_dot_inputs(with_pairs=True) + dot = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot.fit(iterations=5, verbose=False) + assert dot.get_weights().shape[0] == spatial_adata.n_obs + + def test_fit_with_cluster_weight_runs(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, cluster_weight=1.0, device="cpu") + dot.fit(iterations=5, verbose=False) + weights = dot.get_weights() + assert weights.shape[0] == spatial_adata.n_obs + assert np.isfinite(weights.values).all() + + def test_fit_with_ratios_weight_runs(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ratios_weight=1.0, device="cpu") + dot.fit(iterations=5, verbose=False) + weights = dot.get_weights() + assert weights.shape[0] == spatial_adata.n_obs + assert np.isfinite(weights.values).all() + + def test_dict_and_anndata_inputs_produce_matching_weights(self): + ref_adata, spatial_adata = _make_dot_inputs() + ref_dict = _ref_dict_from_adata(ref_adata) + spatial_dict = _spatial_dict_from_adata(spatial_adata) + + dot_adata = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot_adata.fit(iterations=5, verbose=False) + + with pytest.warns(FutureWarning, match="deprecated"): + dot_dict = DOT(spatial=spatial_dict, ref=ref_dict, device="cpu") + dot_dict.fit(iterations=5, verbose=False) + + w_adata = dot_adata.get_weights() + w_dict = dot_dict.get_weights() + assert list(w_adata.columns) == list(w_dict.columns) + assert np.allclose(w_adata.values, w_dict.values) + + +class TestCheckpointing: + def test_save_and_load_checkpoint_round_trips(self, tmp_path): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, ls_solution=False, device="cpu") + C, S = dot.ref.n_obs, dot.spatial.n_obs + + Yt = torch.rand(C, S) + Y_best = torch.rand(C, S) + history = {"iteration": [1, 2], "objective": [10.0, 9.0]} + + dot._save_checkpoint( + str(tmp_path), iteration=2, Yt=Yt, Y_best=Y_best, + f_best=9.0, lb=8.5, history=history, verbose=False, + ) + ckpt_path = tmp_path / "checkpoint_iter_2.pkl" + assert ckpt_path.exists() + + next_iter = dot._load_checkpoint(str(ckpt_path), verbose=False) + assert next_iter == 3 + assert np.allclose(dot.solution, Y_best.numpy()) + assert dot.history == history + + def test_fit_saves_checkpoint_files_at_expected_frequency(self, tmp_path): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot.fit(iterations=6, verbose=False, checkpoint_dir=str(tmp_path), checkpoint_freq=2) + + saved = sorted(p.name for p in tmp_path.glob("checkpoint_iter_*.pkl")) + assert saved == ["checkpoint_iter_2.pkl", "checkpoint_iter_4.pkl", "checkpoint_iter_6.pkl"] + + def test_resume_continues_from_checkpoint_with_comparable_quality(self, tmp_path): + # Resuming restarts from Y_best (not the live Yt at the point of + # interruption) and re-derives Yt through the same warm-start + # rescaling logic used for a fresh ls_solution -- so a resumed run + # is not bit-identical to an uninterrupted one (confirmed + # empirically: weights can differ by ~0.09 in absolute terms even + # though both reach comparable objective quality). Assert on + # objective closeness, not weight equality. + ref_adata, spatial_adata = _make_dot_inputs() + + dot_full = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot_full.fit(iterations=10, verbose=False) + + dot_part = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot_part.fit(iterations=5, verbose=False, checkpoint_dir=str(tmp_path), checkpoint_freq=5) + dot_part.fit(iterations=10, verbose=False, resume_from=str(tmp_path / "checkpoint_iter_5.pkl")) + + obj_full = dot_full.history["objective"][-1] + obj_part = dot_part.history["objective"][-1] + assert obj_part == pytest.approx(obj_full, rel=0.01) + + # history is scoped to the call that produced it, not cumulative + # across the resume boundary -- a real, non-obvious behavior + # worth pinning down rather than assuming continuity. + assert dot_part.history["iteration"] == [6, 7, 8, 9, 10] + + assert np.isfinite(dot_part.get_weights().values).all() + + +def _cluster_wise_mismatch(dot): + """ + Mean cosine distance between each reference sub-cluster's + spatially-implied profile and its own reference profile, computed + independently of DOT's internals -- what l_c/cluster_weight is + supposed to reduce. + """ + Yt = dot.solution + X_sp = dot.spatial.X + X_sp = X_sp.toarray() if issparse(X_sp) else np.asarray(X_sp) + X_ref = dot.ref.X + X_ref = X_ref.toarray() if issparse(X_ref) else np.asarray(X_ref) + + agg = Yt @ X_sp + agg_norms = np.linalg.norm(agg, axis=1, keepdims=True) + agg_norms[agg_norms == 0] = 1 + agg_n = agg / agg_norms + + ref_norms = np.linalg.norm(X_ref, axis=1, keepdims=True) + ref_norms[ref_norms == 0] = 1 + ref_n = X_ref / ref_norms + + cos_sim = (agg_n * ref_n).sum(axis=1) + return float((1 - cos_sim).mean()) + + +class TestClusterWiseCosine: + def test_higher_cluster_weight_reduces_cluster_profile_mismatch(self): + ref_adata, spatial_adata = _make_dot_inputs() + + dot_default = DOT(spatial=spatial_adata, ref=ref_adata, device="cpu") + dot_default.fit(iterations=30, verbose=False) + + dot_cw = DOT(spatial=spatial_adata, ref=ref_adata, cluster_weight=10.0, device="cpu") + dot_cw.fit(iterations=30, verbose=False) + + assert _cluster_wise_mismatch(dot_cw) < _cluster_wise_mismatch(dot_default) + + def test_batched_and_unbatched_cluster_weight_agree(self): + ref_adata, spatial_adata = _make_dot_inputs() + + dot_single_batch = DOT( + spatial=spatial_adata, ref=ref_adata, cluster_weight=5.0, batch_size=1000, device="cpu", + ) + dot_single_batch.fit(iterations=15, verbose=False) + + dot_many_batches = DOT( + spatial=spatial_adata, ref=ref_adata, cluster_weight=5.0, batch_size=1, device="cpu", + ) + dot_many_batches.fit(iterations=15, verbose=False) + + assert np.allclose( + dot_single_batch.get_weights().values, dot_many_batches.get_weights().values, atol=1e-5, + ) + + def test_fit_with_cluster_weight_and_mixed_precision_runs(self): + ref_adata, spatial_adata = _make_dot_inputs() + dot = DOT(spatial=spatial_adata, ref=ref_adata, cluster_weight=1.0, device="cpu") + dot.fit(iterations=5, verbose=False, use_mixed_precision=True) + weights = dot.get_weights() + assert np.isfinite(weights.values).all() diff --git a/tests/test_ds.py b/tests/test_ds.py new file mode 100644 index 0000000..e094812 --- /dev/null +++ b/tests/test_ds.py @@ -0,0 +1,108 @@ +import numpy as np +import pytest + +from dotpy.ds import toy_reference + + +class TestToyReference: + def test_output_shapes_and_dtypes(self): + X, major_types, subtypes = toy_reference( + n_major_types=2, n_subtypes_per_major=3, n_cells_per_subtype=10, + n_genes=50, random_state=0, + ) + n_cells = 2 * 3 * 10 + assert X.shape == (n_cells, 50) + assert major_types.shape == (n_cells,) + assert subtypes.shape == (n_cells,) + assert X.dtype == np.float64 + assert all(isinstance(v, str) for v in major_types) + assert all(isinstance(v, str) for v in subtypes) + + def test_default_naming(self): + _, major_types, subtypes = toy_reference(random_state=0) + assert set(major_types) == {"Major0", "Major1"} + assert set(subtypes) == { + "Major0_Sub0", "Major0_Sub1", "Major0_Sub2", + "Major1_Sub0", "Major1_Sub1", "Major1_Sub2", + } + + def test_custom_naming(self): + _, major_types, subtypes = toy_reference( + n_major_types=2, n_subtypes_per_major=2, + major_prefix="CT", subtype_prefix="State", name_sep=".", + random_state=0, + ) + assert set(major_types) == {"CT0", "CT1"} + assert set(subtypes) == {"CT0.State0", "CT0.State1", "CT1.State0", "CT1.State1"} + + def test_subtype_names_globally_unique(self): + _, _, subtypes = toy_reference( + n_major_types=4, n_subtypes_per_major=5, n_genes=200, random_state=0 + ) + assert len(set(subtypes)) == 4 * 5 + + def test_marker_structure_placement(self): + # Small, hand-verifiable layout: major0=[0:3), sub00=[3:5), sub01=[5:7), + # major1=[7:10), sub10=[10:12), sub11=[12:14), background=[14:20). + # subtype_effect must clear the background's max (9) by more than + # its min (1) undershoots it, i.e. > 8, or an elevated and an + # unelevated gene can land on the same value (e.g. 1+8 == 9). + X, major_types, subtypes = toy_reference( + n_major_types=2, n_subtypes_per_major=2, n_cells_per_subtype=20, + n_genes=20, n_major_markers=3, n_subtype_markers=2, + major_effect=20.0, subtype_effect=15.0, random_state=0, + ) + major0 = major_types == "Major0" + major1 = major_types == "Major1" + sub00 = subtypes == "Major0_Sub0" + sub01 = subtypes == "Major0_Sub1" + + # Major0's marker block: elevated for all Major0 cells, not Major1's. + assert (X[major0][:, 0:3] > 9).all() + assert (X[major1][:, 0:3] <= 9).all() + + # Major0_Sub0's marker block: elevated only for its own cells, not + # even for Major0_Sub1 (same major type, different subtype). + assert (X[sub00][:, 3:5] > 9).all() + assert (X[sub01][:, 3:5] <= 9).all() + + # Pure background genes: never elevated for anyone. + assert (X[:, 14:20] <= 9).all() + + def test_custom_background_range(self): + X, _, _ = toy_reference( + n_major_types=1, n_subtypes_per_major=1, n_genes=50, + n_major_markers=0, n_subtype_markers=0, + background_range=(100, 105), random_state=0, + ) + assert X.min() >= 100 + assert X.max() < 105 + + @pytest.mark.parametrize("background_range", [(5, 5), (5, 1), (1, 2, 3)]) + def test_background_range_invalid_raises(self, background_range): + with pytest.raises(ValueError, match="background_range"): + toy_reference(background_range=background_range) + + def test_n_genes_too_small_raises(self): + with pytest.raises(ValueError, match=r"n_genes=\d+ too small"): + toy_reference(n_major_types=2, n_subtypes_per_major=3, n_genes=5) + + @pytest.mark.parametrize("param_name", ["major_prefix", "subtype_prefix", "name_sep"]) + def test_non_str_naming_param_raises(self, param_name): + with pytest.raises(TypeError, match=param_name): + toy_reference(**{param_name: 123}) + + def test_name_sep_and_subtype_prefix_both_empty_raises(self): + with pytest.raises(ValueError, match="cannot both be empty"): + toy_reference(name_sep="", subtype_prefix="") + + @pytest.mark.parametrize("kwargs", [{"name_sep": "", "subtype_prefix": "S"}, {"name_sep": "-", "subtype_prefix": ""}]) + def test_single_empty_naming_param_ok(self, kwargs): + toy_reference(random_state=0, **kwargs) # no raise + + def test_reproducible_with_fixed_random_state(self): + X1, major1, sub1 = toy_reference(random_state=42) + X2, major2, sub2 = toy_reference(random_state=42) + assert np.array_equal(X1, X2) + assert list(major1) == list(major2) + assert list(sub1) == list(sub2) diff --git a/tests/test_preprocessing.py b/tests/test_preprocessing.py new file mode 100644 index 0000000..b89f6a6 --- /dev/null +++ b/tests/test_preprocessing.py @@ -0,0 +1,2197 @@ +import re + +import numpy as np +import pandas as pd +import pytest +import scanpy as sc +from anndata import AnnData +from scipy.sparse import csr_matrix, issparse + +from dotpy.preprocessing import ( + _capped_group_centroids, + _de_scores_logfc_rank, + _default_k_heuristic, + _expected_k, + _huntington_hill_apportion, + _select_kmeans_genes, + _symmetric_pairs_matrix, + check_subtype_consistency, + kmeans_define_subtypes, + plan_subtype_refinement, + predefined_subtypes, + refine_predefined_subtypes, + select_de_genes, + setup_reference, + setup_spatial, + summarize_subtypes, + validate_reference_input, +) +from _helpers import ( + _make_clustered_spatial_adata, + _make_expression, + _make_obs, + _make_ref_adata, + _make_spatial_adata, +) + + +def _consistent_3level_obs(): + return pd.DataFrame({ + "major": ["A", "A", "A", "A", "B", "B", "B", "B"], + "meso": ["A1", "A1", "A2", "A2", "B1", "B1", "B2", "B2"], + "fine": ["A1a", "A1b", "A2a", "A2b", "B1a", "B1b", "B2a", "B2b"], + }) + + +class TestCheckSubtypeConsistency: + def test_consistent_two_level_passes(self): + obs = _consistent_3level_obs().drop(columns="fine") + check_subtype_consistency(obs, ["major", "meso"]) # no raise + + def test_inconsistent_two_level_raises(self): + obs = pd.DataFrame({ + "major": ["A", "A", "B", "B"], + "subtype": ["S1", "S1", "S1", "S2"], # S1 under both A and B + }) + with pytest.raises(ValueError, match="Inconsistent hierarchy"): + check_subtype_consistency(obs, ["major", "subtype"]) + + def test_consistent_three_level_passes(self): + obs = _consistent_3level_obs() + check_subtype_consistency(obs, ["major", "meso", "fine"]) # no raise + + def test_three_level_pinpoints_offending_pair(self): + obs = pd.DataFrame({ + "major": ["A", "A", "B", "B"], + "meso": ["M1", "M1", "M2", "M2"], # major<->meso: consistent + "fine": ["F1", "F1", "F1", "F2"], # meso<->fine: F1 under M1 and M2 + }) + with pytest.raises(ValueError) as exc_info: + check_subtype_consistency(obs, ["major", "meso", "fine"]) + msg = str(exc_info.value) + assert "'meso'" in msg and "'fine'" in msg + assert "'major'" not in msg + + def test_depth_limits_checked_levels(self): + # Same data as test_three_level_pinpoints_offending_pair: the + # meso<->fine pair is broken, but depth=2 should only check + # major<->meso (which is fine), so this must not raise. + obs = pd.DataFrame({ + "major": ["A", "A", "B", "B"], + "meso": ["M1", "M1", "M2", "M2"], + "fine": ["F1", "F1", "F1", "F2"], + }) + check_subtype_consistency(obs, ["major", "meso", "fine"], depth=2) # no raise + + def test_depth_none_checks_everything(self): + obs = pd.DataFrame({ + "major": ["A", "A", "B", "B"], + "meso": ["M1", "M1", "M2", "M2"], + "fine": ["F1", "F1", "F1", "F2"], + }) + with pytest.raises(ValueError, match="Inconsistent hierarchy"): + check_subtype_consistency(obs, ["major", "meso", "fine"], depth=None) + + @pytest.mark.parametrize("depth", [0, 1, 4]) + def test_depth_out_of_range_raises(self, depth): + obs = _consistent_3level_obs() + with pytest.raises(ValueError, match="depth"): + check_subtype_consistency(obs, ["major", "meso", "fine"], depth=depth) + + def test_fewer_than_two_levels_raises(self): + obs = _consistent_3level_obs() + with pytest.raises(ValueError, match="at least 2"): + check_subtype_consistency(obs, ["major"]) + + def test_missing_column_raises(self): + obs = pd.DataFrame({"major": ["A", "B"]}) + with pytest.raises(ValueError, match="not found in obs.columns"): + check_subtype_consistency(obs, ["major", "subtype"]) + + def test_nan_in_checked_column_raises(self): + obs = pd.DataFrame({ + "major": ["A", "A", None, "B"], + "subtype": ["A1", "A1", "S1", "B1"], + }) + with pytest.raises(ValueError, match="missing \\(NaN\\)"): + check_subtype_consistency(obs, ["major", "subtype"]) + + def test_categorical_dtype_does_not_spuriously_flag_unused_categories(self): + # observed=True must be used so unused categories in a Categorical + # column don't get treated as phantom groups. + obs = pd.DataFrame({ + "major": pd.Categorical(["A", "A", "B", "B"], categories=["A", "B", "C"]), + "subtype": ["A1", "A1", "B1", "B2"], + }) + check_subtype_consistency(obs, ["major", "subtype"]) # no raise + + +class TestExpectedK: + def test_matches_rounded_heuristic_when_within_cap(self): + assert _expected_k(300, 10, _default_k_heuristic) == round(_default_k_heuristic(300)) + + def test_caps_at_subcluster_size(self): + assert _expected_k(10_000_000, 3, _default_k_heuristic) == 3 + + def test_floors_at_one(self): + assert _expected_k(5, 10, _default_k_heuristic) == 1 + + def test_uses_supplied_k_heuristic_not_default(self): + assert _expected_k(100, 10, lambda n: 7.0) == 7 + + def test_rounds_half_to_even(self): + # np.round uses banker's rounding: 2.5 -> 2, 3.5 -> 4. + assert _expected_k(1, 10, lambda n: 2.5) == 2 + assert _expected_k(1, 10, lambda n: 3.5) == 4 + + +class TestHuntingtonHillApportion: + def test_every_group_gets_at_least_one(self): + result = _huntington_hill_apportion({"A": 1000, "B": 1}, 5) + assert result["A"] >= 1 + assert result["B"] >= 1 + + @pytest.mark.parametrize("weights,budget", [ + ({"A": 500, "B": 300, "C": 50, "D": 20}, 8), + ({"A": 10, "B": 10, "C": 10}, 6), + ({"A": 100}, 5), + ]) + def test_budget_is_fully_allocated(self, weights, budget): + result = _huntington_hill_apportion(weights, budget) + assert sum(result.values()) == budget + + def test_no_change_when_budget_equals_group_count(self): + weights = {"A": 500, "B": 300, "C": 50, "D": 20} + result = _huntington_hill_apportion(weights, len(weights)) + assert result == {label: 1 for label in weights} + + def test_does_not_give_all_extra_to_largest(self): + # 4 extra units beyond the guaranteed 1 each; A is by far the + # largest but should not take all 4. + result = _huntington_hill_apportion({"A": 500, "B": 300, "C": 50, "D": 20}, 8) + assert result == {"A": 4, "B": 2, "C": 1, "D": 1} + assert result["A"] < 5 + + def test_single_leftover_unit_goes_to_current_top_priority(self): + # Only 1 extra unit beyond the guaranteed 1 each -- goes to the + # clear-largest group, everyone else stays at their guaranteed 1. + result = _huntington_hill_apportion({"A": 100, "B": 10, "C": 5}, 4) + assert result == {"A": 2, "B": 1, "C": 1} + + def test_single_group_gets_full_budget(self): + assert _huntington_hill_apportion({"A": 100}, 5) == {"A": 5} + + def test_equal_weights_split_evenly(self): + result = _huntington_hill_apportion({"A": 10, "B": 10, "C": 10}, 6) + assert result == {"A": 2, "B": 2, "C": 2} + + def test_larger_weight_never_gets_fewer_seats(self): + result = _huntington_hill_apportion({"A": 100, "B": 50, "C": 25}, 14) + assert result["A"] > result["B"] > result["C"] + + def test_returns_same_keys_as_input(self): + weights = {"A": 500, "B": 300, "C": 50, "D": 20} + result = _huntington_hill_apportion(weights, 8) + assert set(result.keys()) == set(weights.keys()) + + def test_empty_weights_raises(self): + with pytest.raises(ValueError, match="weights must not be empty"): + _huntington_hill_apportion({}, 1) + + @pytest.mark.parametrize("bad_weight", [0, -1]) + def test_nonpositive_weight_raises(self, bad_weight): + with pytest.raises(ValueError, match="must be positive"): + _huntington_hill_apportion({"A": 10, "B": bad_weight}, 3) + + def test_budget_below_group_count_raises(self): + with pytest.raises(ValueError, match="budget"): + _huntington_hill_apportion({"A": 10, "B": 10, "C": 10}, 2) + + +class TestCappedGroupCentroids: + def test_centroid_is_mean_of_the_right_cells(self): + X = np.vstack([np.full((5, 3), 2.0), np.full((5, 3), 10.0)]) + group_labels = np.array(["A"] * 5 + ["B"] * 5) + centroids, _ = _capped_group_centroids(X, group_labels) + # np.unique sorts groups alphabetically: A, B. + assert np.allclose(centroids[0], 2.0) + assert np.allclose(centroids[1], 10.0) + + def test_ratios_reflect_true_relative_sizes(self): + # 3 groups so the ratio values actually matter -- with only 2, + # _select_kmeans_genes's own renormalization would collapse any + # ratio to 1.0, hiding a bug here entirely. + X = _make_expression(100, 5, seed=0) + group_labels = np.array(["A"] * 10 + ["B"] * 30 + ["C"] * 60) + _, ratios = _capped_group_centroids(X, group_labels) + assert ratios["A"] == pytest.approx(0.1) + assert ratios["B"] == pytest.approx(0.3) + assert ratios["C"] == pytest.approx(0.6) + + def test_ratios_use_true_count_not_capped_sample_count(self): + # Both groups exceed cap=10, so both are subsampled down to the + # same size -- if ratios were derived from the sample instead of + # the true group size, they'd come out equal (10/10) instead of + # reflecting the true 50/100 split. + X = _make_expression(150, 5, seed=1) + group_labels = np.array(["A"] * 50 + ["B"] * 100) + _, ratios = _capped_group_centroids(X, group_labels, cap=10) + assert ratios["A"] == pytest.approx(50 / 150) + assert ratios["B"] == pytest.approx(100 / 150) + + def test_cap_bounds_cells_used_for_centroid(self): + # Distinct, spread-out per-cell values so a partial-sample mean is + # essentially guaranteed to differ from the true full-population + # mean, for a fixed seed. + X = np.arange(2000 * 3, dtype=float).reshape(2000, 3) + group_labels = np.array(["A"] * 2000) + + np.random.seed(0) + centroid_capped, _ = _capped_group_centroids(X, group_labels, cap=50) + np.random.seed(0) + centroid_full, _ = _capped_group_centroids(X, group_labels, cap=2000) # == true size, no subsampling + + assert not np.allclose(centroid_capped, centroid_full) + + def test_reproducible_given_fixed_ambient_random_state(self): + X = _make_expression(200, 5, seed=2) + group_labels = np.array(["A"] * 150 + ["B"] * 50) + + np.random.seed(0) + centroids_1, ratios_1 = _capped_group_centroids(X, group_labels, cap=20) + np.random.seed(0) + centroids_2, ratios_2 = _capped_group_centroids(X, group_labels, cap=20) + + assert np.array_equal(centroids_1, centroids_2) + assert ratios_1 == ratios_2 + + def test_sparse_and_dense_input_agree(self): + X_dense = _make_expression(50, 5, seed=3) + X_sparse = csr_matrix(X_dense) + group_labels = np.array(["A"] * 25 + ["B"] * 25) + + centroids_dense, ratios_dense = _capped_group_centroids(X_dense, group_labels) + centroids_sparse, ratios_sparse = _capped_group_centroids(X_sparse, group_labels) + + assert np.allclose(centroids_dense, centroids_sparse) + assert ratios_dense == ratios_sparse + + +class TestSelectKmeansGenes: + def test_single_major_type_returns_all_genes(self): + # No "other" major types to compare against -- can't compute a + # log-fold change, so every gene is returned unfiltered. + ct_centroid = np.array([1.0, 2.0, 3.0]) + major_centroids = np.array([[1.0, 2.0, 3.0]]) + result = _select_kmeans_genes(ct_centroid, major_centroids, {"A": 1.0}, "A") + assert list(result) == [0, 1, 2] + + def test_more_passing_genes_than_max_genes_truncates_to_top_logfc(self): + ct_centroid = np.array([100.0, 90.0, 80.0, 70.0, 60.0]) + major_centroids = np.array([ + [0.0, 0.0, 0.0, 0.0, 0.0], # A's own row, unused (ct_centroid is separate) + [1.0, 1.0, 1.0, 1.0, 1.0], # B's centroid + ]) + result = _select_kmeans_genes( + ct_centroid, major_centroids, {"A": 1.0, "B": 1.0}, "A", max_genes=3, + ) + assert list(result) == [0, 1, 2] + + +class TestKmeansDefineSubtypes: + def test_output_shape_and_name_format(self): + X = _make_expression(60, 20, seed=0) + annotations = np.array(["A"] * 30 + ["B"] * 30) + df = kmeans_define_subtypes( + X, annotations, subcluster_size=3, random_state=0, + annotations_key="major", subtype_key="subtype", + ) + assert df.shape == (60, 2) + assert list(df.columns) == ["major", "subtype"] + assert list(df["major"]) == list(annotations) + for lab in df["subtype"]: + if lab is not None: + assert isinstance(lab, str) + major, idx = lab.rsplit("_", 1) + assert major in ("A", "B") + assert idx.isdigit() + + def test_custom_annotations_and_subtype_keys(self): + X = _make_expression(60, 20, seed=20) + annotations = np.array(["A"] * 30 + ["B"] * 30) + df = kmeans_define_subtypes( + X, annotations, subcluster_size=3, random_state=0, + annotations_key="meso", subtype_key="fine", + ) + assert list(df.columns) == ["meso", "fine"] + assert list(df["meso"]) == list(annotations) + + @pytest.mark.parametrize("key", ["major", "custom"]) + def test_annotations_key_equals_subtype_key_raises(self, key): + X = _make_expression(60, 20, seed=21) + annotations = np.array(["A"] * 60) + with pytest.raises(ValueError, match="annotations_key"): + kmeans_define_subtypes(X, annotations, annotations_key=key, subtype_key=key, random_state=0) + + def test_min_frac_zero_drops_nothing(self): + X = _make_expression(300, 20, seed=1) + annotations = np.array(["A"] * 300) + df = kmeans_define_subtypes(X, annotations, random_state=0, min_frac=0.0, max_cells_per_type=1000) + assert all(lab is not None for lab in df["subtype"]) + + def test_min_frac_default_drops_small_distinct_subgroup(self): + # 5 cells with a clearly shifted profile (1.67% of 300, under the + # explicit min_frac=0.025) among 295 similar "background" cells. + # Forcing K=2 makes the shifted group its own cluster, which + # should then be dropped as noise while the background survives. + background = _make_expression(295, 20, seed=12) + rare = _make_expression(5, 20, seed=13) + rare[:, :5] += 50 # shift a block of genes so it's clearly separable + X = np.vstack([background, rare]) + annotations = np.array(["A"] * 300) + with pytest.warns(UserWarning, match=r"Dropped \d+ cell"): + df = kmeans_define_subtypes( + X, annotations, random_state=0, k_heuristic=lambda n: 2, + min_frac=0.025, max_cells_per_type=1000, + ) + subtype = df["subtype"] + assert sum(lab is not None for lab in subtype[:295]) > 250 + assert all(lab is None for lab in subtype[295:]) + + def test_subcluster_size_caps_distinct_subtypes(self): + # Natural K for n=1000 is ~7 (2*log(1000)-7); subcluster_size=3 must cap it. + X = _make_expression(1000, 20, seed=2) + annotations = np.array(["A"] * 1000) + df = kmeans_define_subtypes(X, annotations, subcluster_size=3, random_state=0) + distinct = {lab for lab in df["subtype"] if lab is not None} + assert len(distinct) <= 3 + + def test_k_heuristic_override_changes_subtype_count(self): + # min_frac=0 isolates the k_heuristic wiring itself: without it, + # forcing K=6 on unstructured data could legitimately have some of + # those 6 clusters fall below min_frac and get dropped. + X = _make_expression(300, 20, seed=3) + annotations = np.array(["A"] * 300) + subcluster_size = 10 + + df_default = kmeans_define_subtypes( + X, annotations, random_state=0, min_frac=0.0, subcluster_size=subcluster_size, + max_cells_per_type=1000, + ) + n_default = len({lab for lab in df_default["subtype"] if lab is not None}) + expected_n_default = min(subcluster_size, max(1, int(np.round(_default_k_heuristic(300))))) + assert n_default == expected_n_default + + df_fixed = kmeans_define_subtypes( + X, annotations, random_state=0, min_frac=0.0, k_heuristic=lambda n: 6, + max_cells_per_type=1000, + ) + n_fixed = len({lab for lab in df_fixed["subtype"] if lab is not None}) + assert n_fixed == 6 + + def test_max_cells_per_type_caps_considered_cells(self): + X = _make_expression(500, 20, seed=4) + annotations = np.array(["A"] * 500) + df = kmeans_define_subtypes(X, annotations, random_state=0, max_cells_per_type=50) + n_assigned = int(sum(lab is not None for lab in df["subtype"])) + assert n_assigned <= 50 + + def test_min_frac_uses_sampled_count_not_full_major_type_size(self): + # 2000 true cells, capped to 200 sampled. min_frac=0.025 should + # give a noise floor of 0.025*200=5 (computed from what was + # actually clustered) -- not 0.025*2000=50 (the full major type). + # The ~20-cell rare cluster clears 5 comfortably but would be + # dropped as noise under the much larger, wrong threshold, so + # nothing being dropped here pins the correct denominator. + background = _make_expression(1800, 20, seed=0) + rare = _make_expression(200, 20, seed=100) + rare[:, :5] += 50 # shift a block of genes so it's clearly separable + X = np.vstack([background, rare]) + annotations = np.array(["A"] * 2000) + + df = kmeans_define_subtypes( + X, annotations, random_state=0, k_heuristic=lambda n: 2, + min_frac=0.025, max_cells_per_type=200, + ) + assert int(df["subtype"].notna().sum()) == 200 + assert set(df["subtype"].dropna()) == {"A_0", "A_1"} + + def test_custom_name_sep(self): + X = _make_expression(300, 20, seed=5) + annotations = np.array(["A"] * 300) + df = kmeans_define_subtypes(X, annotations, random_state=0, name_sep="::", min_frac=0.0) + non_none = [lab for lab in df["subtype"] if lab is not None] + assert len(non_none) > 0 + assert all(lab.startswith("A::") for lab in non_none) + + def test_single_cell_major_type_warns_and_unassigned(self): + X = _make_expression(101, 20, seed=6) + annotations = np.array(["A"] * 100 + ["B"] * 1) + with pytest.warns(UserWarning, match="only 1 cell"): + df = kmeans_define_subtypes(X, annotations, random_state=0) + assert df["subtype"].iloc[-1] is None + + def test_extreme_min_frac_forces_single_fallback_subtype(self): + # K>=2 (forced via k_heuristic) means no single cluster can hold + # >=99% of cells, so min_frac=0.99 reliably drops everything as + # noise regardless of the actual data. + X = _make_expression(300, 20, seed=7) + annotations = np.array(["A"] * 300) + with pytest.warns(UserWarning, match="dropped as noise"): + df = kmeans_define_subtypes( + X, annotations, random_state=0, min_frac=0.99, k_heuristic=lambda n: 5, + subcluster_size=10, max_cells_per_type=1000, + ) + non_none = [lab for lab in df["subtype"] if lab is not None] + assert len(non_none) == 300 + assert set(non_none) == {"A_0"} + + def test_few_cells_get_single_subtype_no_clustering(self): + # k_heuristic pinned to guarantee K<=1 (the shortcut under test) + # regardless of what the default formula returns for 5 cells; this + # also makes subcluster_size/min_frac/max_cells_per_type irrelevant + # here, since the shortcut skips the clustering path entirely. + X = _make_expression(5, 20, seed=8) + annotations = np.array(["A"] * 5) + df = kmeans_define_subtypes(X, annotations, random_state=0, k_heuristic=lambda n: 1) + assert list(df["subtype"]) == ["A_0"] * 5 + + def test_reproducible_with_fixed_random_state(self): + X = _make_expression(300, 20, seed=9) + annotations = np.array(["A"] * 150 + ["B"] * 150) + df1 = kmeans_define_subtypes(X, annotations, random_state=42) + df2 = kmeans_define_subtypes(X, annotations, random_state=42) + pd.testing.assert_frame_equal(df1, df2) + + def test_sparse_and_dense_input_agree(self): + X_dense = _make_expression(300, 20, seed=10) + X_sparse = csr_matrix(X_dense) + annotations = np.array(["A"] * 150 + ["B"] * 150) + df_dense = kmeans_define_subtypes(X_dense, annotations, random_state=0) + df_sparse = kmeans_define_subtypes(X_sparse, annotations, random_state=0) + pd.testing.assert_frame_equal(df_dense, df_sparse) + + def test_gene_selection_path_with_many_genes(self): + # >500 genes routes through _select_kmeans_genes; two major types + # so it actually compares one against the other, not the trivial + # single-type shortcut. + X = _make_expression(200, 600, seed=11) + annotations = np.array(["A"] * 100 + ["B"] * 100) + df = kmeans_define_subtypes(X, annotations, random_state=0) + assert len(df) == 200 + assert any(lab is not None for lab in df["subtype"]) + + def test_th_inner_logfold_reaches_gene_selection(self, capsys): + X_a = _make_expression(100, 600, seed=13) + X_a[:, :50] += 20 + X_b = _make_expression(100, 600, seed=14) + X = np.vstack([X_a, X_b]) + annotations = np.array(["A"] * 100 + ["B"] * 100) + + kmeans_define_subtypes(X, annotations, random_state=0, th_inner_logfold=0.5, verbose=True) + out_loose = capsys.readouterr().out + kmeans_define_subtypes(X, annotations, random_state=0, th_inner_logfold=2.0, verbose=True) + out_strict = capsys.readouterr().out + + # major_types is alphabetically sorted, so A's line prints first. + n_loose = int(re.search(r"using (\d+) genes", out_loose).group(1)) + n_strict = int(re.search(r"using (\d+) genes", out_strict).group(1)) + assert n_loose == 50 + assert n_strict == 600 + + def test_multiple_major_types_hit_different_branches_in_one_call(self): + # A: enough cells for normal clustering. B: exactly 1 cell (skipped + # with a warning). C: few cells, yields K<=1. k_heuristic is pinned + # to guarantee this branch split (K=5 for A, K=1 for C) regardless + # of what the default formula would return for these cell counts. + X_a = _make_expression(300, 20, seed=14) + X_b = _make_expression(1, 20, seed=15) + X_c = _make_expression(5, 20, seed=16) + X = np.vstack([X_a, X_b, X_c]) + annotations = np.array(["A"] * 300 + ["B"] * 1 + ["C"] * 5) + + with pytest.warns(UserWarning, match="only 1 cell"): + df = kmeans_define_subtypes( + X, annotations, random_state=0, + k_heuristic=lambda n: 5 if n > 100 else 1, + subcluster_size=10, max_cells_per_type=1000, + ) + + subtype = df["subtype"] + assert any(lab is not None and lab.startswith("A_") for lab in subtype[:300]) + assert subtype.iloc[300] is None + assert list(subtype[301:]) == ["C_0"] * 5 + + def test_all_major_types_too_small_returns_all_none(self): + X = _make_expression(2, 20, seed=17) + annotations = np.array(["A", "B"]) + with pytest.warns(UserWarning): + df = kmeans_define_subtypes(X, annotations, random_state=0) + assert list(df["subtype"]) == [None, None] + + @pytest.mark.parametrize("subcluster_size", [0, -1]) + def test_subcluster_size_below_one_raises(self, subcluster_size): + X = _make_expression(60, 20, seed=18) + annotations = np.array(["A"] * 60) + with pytest.raises(ValueError, match="subcluster_size"): + kmeans_define_subtypes(X, annotations, subcluster_size=subcluster_size, random_state=0) + + @pytest.mark.parametrize("max_cells_per_type", [0, -1]) + def test_max_cells_per_type_below_one_raises(self, max_cells_per_type): + X = _make_expression(60, 20, seed=19) + annotations = np.array(["A"] * 60) + with pytest.raises(ValueError, match="max_cells_per_type"): + kmeans_define_subtypes(X, annotations, max_cells_per_type=max_cells_per_type, random_state=0) + + +class TestPlanSubtypeRefinement: + def test_return_counts_false_returns_none(self): + # k_heuristic matches the 2 labels exactly -- no incidental warning. + obs = _make_obs({"A": [500, 500]}) + assert plan_subtype_refinement(obs, ["major", "subtype"], k_heuristic=lambda n: 2, subcluster_size=10) is None + + def test_return_counts_true_returns_dataframe_with_expected_columns(self): + # k_heuristic matches each major type's own label count exactly + # (A: 1000 cells/2 labels, B: 900 cells/3 labels) -- no incidental + # warning. + obs = _make_obs({"A": [500, 500], "B": [300, 300, 300]}) + result = plan_subtype_refinement( + obs, ["major", "subtype"], k_heuristic=lambda n: 2 if n == 1000 else 3, return_counts=True, + ) + assert list(result.columns) == [ + "major", "subtype", "n_cells", "major_type_fraction", "current_k", "expected_k", "budget", + ] + assert len(result) == 5 # 2 + 3 distinct subtype labels + + def test_no_warning_when_labels_match_expected_k(self, recwarn): + # k_heuristic pinned to 2 so expected_k == the actual 2 labels + # exactly -- neither over- nor under-clustered -- and 500/500 is + # well above min_frac, so no small-label warning either. + obs = _make_obs({"A": [500, 500]}) + plan_subtype_refinement(obs, ["major", "subtype"], k_heuristic=lambda n: 2, subcluster_size=10) + assert len(recwarn) == 0 + + def test_overclustering_warns(self): + obs = _make_obs({"A": [100] * 10}) # 10 subtypes, 1000 cells + with pytest.warns(UserWarning, match=r"more '.*' labels"): + plan_subtype_refinement(obs, ["major", "subtype"], subcluster_size=10, k_heuristic=_default_k_heuristic) + + def test_underclustering_warns(self): + obs = _make_obs({"A": [1000]}) # 1 subtype only + with pytest.warns(UserWarning, match=r"fewer '.*' labels"): + plan_subtype_refinement(obs, ["major", "subtype"], k_heuristic=lambda n: 5, subcluster_size=10) + + def test_min_frac_warns_for_undersized_label(self): + # k_heuristic pinned to 2 so expected_k matches the 2 labels + # exactly -- isolates the min_frac warning from the + # over-/under-clustering ones. + obs = _make_obs({"A": [990, 10]}) # 10/1000 = 1%, under min_frac + with pytest.warns(UserWarning, match=r"below min_frac"): + plan_subtype_refinement(obs, ["major", "subtype"], min_frac=0.025, k_heuristic=lambda n: 2) + + def test_multiple_major_types_evaluated_independently(self): + obs_a = _make_obs({"A": [100] * 10}) # 10 labels, over-clustered vs expected_k=5 + obs_b = _make_obs({"B": [200] * 5}) # 5 labels, matches expected_k=5 exactly + obs = pd.concat([obs_a, obs_b], ignore_index=True) + with pytest.warns(UserWarning) as record: + plan_subtype_refinement( + obs, ["major", "subtype"], + subcluster_size=5, k_heuristic=lambda n: 5, min_frac=0.0, + ) + messages = [str(w.message) for w in record] + assert any("'A'" in m for m in messages) + assert not any("'B'" in m for m in messages) + + def test_multiple_overclustered_major_types_aggregate_into_one_warning(self): + obs_a = _make_obs({"A": [100] * 10}) # over-clustered + obs_b = _make_obs({"B": [100] * 8}) # over-clustered + obs_c = _make_obs({"C": [100] * 5}) # matches expected_k=5 exactly + obs = pd.concat([obs_a, obs_b, obs_c], ignore_index=True) + with pytest.warns(UserWarning) as record: + plan_subtype_refinement( + obs, ["major", "subtype"], subcluster_size=10, k_heuristic=lambda n: 5, min_frac=0.0, + ) + messages = [str(w.message) for w in record] + assert len(messages) == 1 + assert "['A', 'B']" in messages[0] + + def test_multiple_underclustered_major_types_aggregate_into_one_warning(self): + obs_d = _make_obs({"D": [1000]}) # under-clustered + obs_e = _make_obs({"E": [500, 500]}) # under-clustered + obs_f = _make_obs({"F": [200] * 5}) # matches expected_k=5 exactly + obs = pd.concat([obs_d, obs_e, obs_f], ignore_index=True) + with pytest.warns(UserWarning) as record: + plan_subtype_refinement( + obs, ["major", "subtype"], subcluster_size=10, k_heuristic=lambda n: 5, min_frac=0.0, + ) + messages = [str(w.message) for w in record] + assert len(messages) == 1 + assert "['D', 'E']" in messages[0] + + def test_min_frac_warning_lists_undersized_count_over_total_per_major_type(self): + obs_g = _make_obs({"G": [990, 10]}) # 1/2 labels undersized + obs_h = _make_obs({"H": [970, 10, 10, 10]}) # 3/4 labels undersized + obs = pd.concat([obs_g, obs_h], ignore_index=True) + with pytest.warns(UserWarning) as record: + plan_subtype_refinement( + obs, ["major", "subtype"], + k_heuristic=lambda n: 2 if n == 1000 else 4, subcluster_size=10, min_frac=0.025, + ) + messages = [str(w.message) for w in record] + assert any("G: 1/2" in m and "H: 3/4" in m for m in messages) + + def test_all_three_warning_categories_fire_independently_in_one_call(self): + obs_i = _make_obs({"I": [100] * 10}) # over-clustered + obs_j = _make_obs({"J": [1000]}) # under-clustered + obs_k = _make_obs({"K": [990, 10]}) # under-clustered and has an undersized label + obs = pd.concat([obs_i, obs_j, obs_k], ignore_index=True) + with pytest.warns(UserWarning) as record: + plan_subtype_refinement( + obs, ["major", "subtype"], subcluster_size=10, k_heuristic=lambda n: 5, min_frac=0.025, + ) + assert len(record) == 3 + + def test_major_type_fraction_matches_label_share_of_major_type_total(self): + obs = _make_obs({"A": [250, 750]}) + plan = plan_subtype_refinement( + obs, ["major", "subtype"], k_heuristic=lambda n: 2, subcluster_size=10, return_counts=True, + ).set_index("subtype") + assert plan.loc["A_0", "major_type_fraction"] == pytest.approx(0.25) + assert plan.loc["A_1", "major_type_fraction"] == pytest.approx(0.75) + + def test_budget_and_n_cells_are_int_with_categorical_labels(self): + # AnnData obs columns are categorical by default. pandas' Series.map() + # silently returns float/category dtype (not int) on a categorical + # Series even when every value matches a dict key -- a real bug this + # pins, since a float budget crashes sklearn's KMeans(n_clusters=...). + obs = _make_obs({"A": [400, 100]}) + obs["major"] = obs["major"].astype("category") + obs["subtype"] = obs["subtype"].astype("category") + with pytest.warns(UserWarning, match=r"fewer '.*' labels"): + plan = plan_subtype_refinement( + obs, ["major", "subtype"], k_heuristic=lambda n: 5, subcluster_size=10, + min_frac=0.0, return_counts=True, + ) + assert plan["n_cells"].dtype == np.int64 + assert plan["budget"].dtype == np.int64 + + def test_only_finest_level_checked_against_coarsest(self): + obs = pd.DataFrame({ + "major": ["T"] * 1000, + "meso": sum([[f"T_m{i}"] * 100 for i in range(10)], []), + "fine": sum([[f"T_f{i}"] * 100 for i in range(10)], []), + }) + with pytest.warns(UserWarning) as record: + plan_subtype_refinement(obs, ["major", "meso", "fine"], subcluster_size=10, k_heuristic=lambda n: 5) + messages = [str(w.message) for w in record] + assert len(messages) == 1 + assert "'fine'" in messages[0] + assert "'meso'" not in messages[0] + + def test_inconsistent_hierarchy_raises(self): + obs = pd.DataFrame({"major": ["A", "B"], "subtype": ["S1", "S1"]}) + with pytest.raises(ValueError, match="Inconsistent hierarchy"): + plan_subtype_refinement(obs, ["major", "subtype"]) + + @pytest.mark.parametrize("subcluster_size", [0, -1]) + def test_subcluster_size_below_one_raises(self, subcluster_size): + obs = _make_obs({"A": [10, 10]}) + with pytest.raises(ValueError, match="subcluster_size"): + plan_subtype_refinement(obs, ["major", "subtype"], subcluster_size=subcluster_size) + + def test_budget_sums_to_max_expected_k_and_current_k_per_major_type(self): + obs_a = _make_obs({"A": [500, 500]}) # under-clustered: 2 labels < expected_k=5 + obs_b = _make_obs({"B": [100] * 10}) # over-clustered: 10 labels > expected_k=5 + obs = pd.concat([obs_a, obs_b], ignore_index=True) + # Both directions are intrinsic to this test (it verifies budget + # sums for one under- and one over-clustered major type at once), + # so both warnings are expected here, not incidental. + with pytest.warns(UserWarning) as record: + plan = plan_subtype_refinement( + obs, ["major", "subtype"], + k_heuristic=lambda n: 5, subcluster_size=10, min_frac=0.0, return_counts=True, + ) + assert len(record) == 2 + sums = plan.groupby("major")["budget"].sum() + assert sums["A"] == 5 + assert sums["B"] == 10 + + def test_not_underclustered_major_type_has_budget_one_everywhere(self): + # k_heuristic matches the 2 labels exactly, so this is neither + # over- nor under-clustered -- no incidental warning. + obs = _make_obs({"A": [500, 500]}) + plan = plan_subtype_refinement( + obs, ["major", "subtype"], k_heuristic=lambda n: 2, subcluster_size=10, return_counts=True, + ) + assert (plan["budget"] == 1).all() + + def test_underclustered_major_type_apportions_by_size(self): + # Under-clustered by design -- this is exactly what the test + # verifies apportionment for, so the warning is expected, not + # incidental. + obs = _make_obs({"A": [800, 200]}) + with pytest.warns(UserWarning, match=r"fewer '.*' labels"): + plan = plan_subtype_refinement( + obs, ["major", "subtype"], + k_heuristic=lambda n: 6, subcluster_size=10, return_counts=True, + ).set_index("subtype") + assert plan["budget"].sum() == 6 + assert (plan["budget"] >= 1).all() + assert plan.loc["A_0", "budget"] > plan.loc["A_1", "budget"] + + def test_three_level_hierarchy_preserves_all_level_keys_in_output(self): + # Both major types have 4 cells and 4 distinct fine labels. + obs = _consistent_3level_obs() + plan = plan_subtype_refinement(obs, ["major", "meso", "fine"], k_heuristic=lambda n: 4, return_counts=True) + assert list(plan.columns[:3]) == ["major", "meso", "fine"] + plan = plan.set_index("fine") + assert plan.loc["A1a", "meso"] == "A1" + assert plan.loc["B2b", "meso"] == "B2" + assert plan.loc["A1a", "major"] == "A" + assert plan.loc["B2b", "major"] == "B" + + def test_n_cells_matches_actual_label_sizes(self): + obs = _make_obs({"A": [30, 70]}) + plan = plan_subtype_refinement(obs, ["major", "subtype"], return_counts=True).set_index("subtype") + assert plan.loc["A_0", "n_cells"] == 30 + assert plan.loc["A_1", "n_cells"] == 70 + + def test_current_k_and_expected_k_repeat_correctly_within_major_type(self): + # Deliberately mismatched (current_k=3, expected_k=7) to verify + # the two columns are tracked independently -- so this is + # necessarily under-clustered, and the warning is expected. + obs = _make_obs({"A": [100, 100, 100]}) # 3 distinct labels + with pytest.warns(UserWarning, match=r"fewer '.*' labels"): + plan = plan_subtype_refinement( + obs, ["major", "subtype"], + k_heuristic=lambda n: 7, subcluster_size=10, return_counts=True, + ) + assert (plan["current_k"] == 3).all() + assert (plan["expected_k"] == 7).all() + + def test_verbose_prints_summary(self, capsys): + # k_heuristic matches the 2 labels exactly, so this is neither + # over- nor under-clustered -- no incidental warning. + obs = _make_obs({"A": [500, 500]}) + plan_subtype_refinement(obs, ["major", "subtype"], k_heuristic=lambda n: 2, subcluster_size=10, verbose=True) + out = capsys.readouterr().out + assert "'major'='A'" in out + assert "'subtype'" in out + + +class TestPredefinedSubtypes: + def test_output_matches_input_labels_and_order(self): + obs = pd.DataFrame({ + "major": ["A", "B", "A", "B"], + "subtype": ["CD8_Naive", "Bmem", "CD8_Other", "Bnaive"], + }) + # min_frac=0.0 and a k_heuristic matching each major type's 2 + # labels exactly make warnings structurally impossible; this test + # only checks passthrough. + result = predefined_subtypes(obs, ["major", "subtype"], min_frac=0.0, k_heuristic=lambda n: 2) + pd.testing.assert_frame_equal(result, obs[["major", "subtype"]]) + + def test_returns_all_levels_for_n_greater_than_two(self): + obs = _consistent_3level_obs() + # Both major types have 4 cells and 4 distinct fine labels. + result = predefined_subtypes(obs, ["major", "meso", "fine"], min_frac=0.0, k_heuristic=lambda n: 4) + pd.testing.assert_frame_equal(result, obs[["major", "meso", "fine"]]) + + def test_categorical_subtype_returns_values_not_codes(self): + obs = pd.DataFrame({ + "major": ["A", "A", "B", "B"], + "subtype": pd.Categorical(["A1", "A1", "B1", "B1"]), + }) + # Both major types have 2 cells and 1 distinct label. + result = predefined_subtypes(obs, ["major", "subtype"], min_frac=0.0, k_heuristic=lambda n: 1) + assert list(result["subtype"]) == ["A1", "A1", "B1", "B1"] + + def test_inconsistent_hierarchy_raises(self): + obs = pd.DataFrame({ + "major": ["A", "B"], + "subtype": ["S1", "S1"], + }) + with pytest.raises(ValueError, match="Inconsistent hierarchy"): + predefined_subtypes(obs, ["major", "subtype"]) + + @pytest.mark.parametrize("subcluster_size", [0, -1]) + def test_subcluster_size_below_one_raises(self, subcluster_size): + obs = _make_obs({"A": [10, 10]}) + with pytest.raises(ValueError, match="subcluster_size"): + predefined_subtypes(obs, ["major", "subtype"], subcluster_size=subcluster_size) + + def test_no_warning_within_expected_k_and_min_frac(self, recwarn): + subcluster_size = 10 + min_frac = 0.025 + k_heuristic = _default_k_heuristic + # 2 balanced subtypes over 100 cells matches expected_k=2 exactly + # under the real default heuristic (neither over- nor + # under-clustered), and both are well above min_frac. + obs = _make_obs({"A": [50, 50]}) + predefined_subtypes( + obs, ["major", "subtype"], + subcluster_size=subcluster_size, min_frac=min_frac, k_heuristic=k_heuristic, + ) + assert len(recwarn) == 0 + + def test_overclustering_warns_when_count_exceeds_expected_k(self): + obs = _make_obs({"A": [100] * 10}) # 10 subtypes, 1000 cells + with pytest.warns(UserWarning, match=r"more '.*' labels"): + predefined_subtypes(obs, ["major", "subtype"], subcluster_size=10, k_heuristic=_default_k_heuristic) + + def test_underclustering_warns_when_count_below_expected_k(self): + obs = _make_obs({"A": [1000]}) # 1 subtype, well under expected_k + with pytest.warns(UserWarning, match=r"fewer '.*' labels"): + predefined_subtypes(obs, ["major", "subtype"], subcluster_size=10, k_heuristic=_default_k_heuristic) + + def test_subcluster_size_caps_expected_k_below_heuristic(self): + # k_heuristic fixed to 7, well above subcluster_size=3, so 3 is + # guaranteed to be the binding cap. The exact expected_k is no + # longer observable via the warning text (only via + # plan_subtype_refinement's return_counts=True), so this only + # confirms subcluster_size capping still triggers over-clustering. + obs = _make_obs({"A": [300, 300, 300, 300]}) + with pytest.warns(UserWarning, match=r"more '.*' labels"): + predefined_subtypes(obs, ["major", "subtype"], subcluster_size=3, k_heuristic=lambda n: 7) + + def test_k_heuristic_override_changes_expected_k(self): + # 400 cells: the real default heuristic gives expected_k~5, so 4 + # labels would NOT over-cluster under it. Forcing k_heuristic to a + # constant 3 makes it over-cluster instead, confirming the + # override is actually used, not silently ignored. + obs = _make_obs({"A": [100] * 4}) + with pytest.warns(UserWarning, match=r"more '.*' labels"): + predefined_subtypes(obs, ["major", "subtype"], k_heuristic=lambda n: 3, subcluster_size=10) + + def test_default_k_heuristic_wiring_matches_formula(self): + # k_heuristic omitted entirely, relying on predefined_subtypes's + # own default parameter binding to _default_k_heuristic. + obs = _make_obs({"A": [50] * 20}) + with pytest.warns(UserWarning, match=r"more '.*' labels"): + predefined_subtypes(obs, ["major", "subtype"], subcluster_size=10) + + def test_min_frac_warns_for_undersized_subtype(self): + min_frac = 0.025 + # k_heuristic pinned to 2 so expected_k matches the 2 labels + # exactly, isolating the min_frac warning. + obs = _make_obs({"A": [990, 10]}) # 10/1000 = 1%, under min_frac + with pytest.warns(UserWarning, match=r"below min_frac"): + predefined_subtypes(obs, ["major", "subtype"], min_frac=min_frac, k_heuristic=lambda n: 2) + + def test_min_frac_zero_disables_size_warning(self, recwarn): + subcluster_size = 10 + k_heuristic = _default_k_heuristic + # 100 total cells matches expected_k=2 exactly for these 2 labels, + # clearing the over-/under-clustering warnings. The 2-cell label is + # 2% of the total, genuinely below the real default min_frac=0.025 + # (2.5%) -- so this only passes because min_frac=0.0 is actually + # doing something, not because nothing would have warned anyway. + obs = _make_obs({"A": [98, 2]}) + predefined_subtypes( + obs, ["major", "subtype"], + min_frac=0.0, subcluster_size=subcluster_size, k_heuristic=k_heuristic, + ) + assert len(recwarn) == 0 + + def test_multiple_major_types_evaluated_independently(self): + obs_a = _make_obs({"A": [100] * 10}) # 10 labels, over-clustered vs expected_k=5 + obs_b = _make_obs({"B": [200] * 5}) # 5 labels, matches expected_k=5 exactly + obs = pd.concat([obs_a, obs_b], ignore_index=True) + # Fixed subcluster_size/k_heuristic give expected_k=5 for both major + # types regardless of cell count: A (10 labels) must warn, B (5) + # must not. min_frac=0.0 keeps the size warning fully out of scope. + with pytest.warns(UserWarning) as record: + predefined_subtypes( + obs, ["major", "subtype"], + subcluster_size=5, k_heuristic=lambda n: 5, min_frac=0.0, + ) + messages = [str(w.message) for w in record] + assert any("'A'" in m for m in messages) + assert not any("'B'" in m for m in messages) + + def test_only_finest_level_checked_against_coarsest(self): + # meso and fine are the same 10-way partition under "T" (fine + # coincides with meso 1:1), so both would independently exceed + # expected_k=5 if checked. Only the finest level should actually + # be examined: exactly one warning, naming only "fine". + obs = pd.DataFrame({ + "major": ["T"] * 1000, + "meso": sum([[f"T_m{i}"] * 100 for i in range(10)], []), + "fine": sum([[f"T_f{i}"] * 100 for i in range(10)], []), + }) + with pytest.warns(UserWarning) as record: + predefined_subtypes(obs, ["major", "meso", "fine"], subcluster_size=10, k_heuristic=lambda n: 5) + messages = [str(w.message) for w in record] + assert len(messages) == 1 + assert "'fine'" in messages[0] + assert "'meso'" not in messages[0] + + def test_verbose_prints_summary(self, capsys): + # 100 total cells matches expected_k=2 exactly for these 2 labels, + # avoiding an incidental under-clustering warning. + obs = _make_obs({"A": [50, 50]}) + predefined_subtypes(obs, ["major", "subtype"], verbose=True) + out = capsys.readouterr().out + assert "'major'='A'" in out + assert "'subtype'" in out + assert "2 labels" in out + assert "expected" in out + + def test_verbose_false_prints_nothing(self, capsys): + obs = _make_obs({"A": [50, 50]}) + predefined_subtypes(obs, ["major", "subtype"], verbose=False) + out = capsys.readouterr().out + assert out == "" + + +class TestRefinePredefinedSubtypes: + def test_passthrough_when_nothing_flagged(self): + obs = _make_obs({"A": [50, 50], "B": [30, 30, 30]}) + X = _make_expression(len(obs), 10, seed=0) + plan = pd.DataFrame({ + "major": ["A", "A", "B", "B", "B"], + "subtype": ["A_0", "A_1", "B_0", "B_1", "B_2"], + "n_cells": [50, 50, 30, 30, 30], + "current_k": [2, 2, 3, 3, 3], + "expected_k": [2, 2, 3, 3, 3], + "budget": [1, 1, 1, 1, 1], + }) + result = refine_predefined_subtypes(X, obs, ["major", "subtype"], plan) + assert list(result.columns) == ["major", "subtype", "subtype_kmeans"] + assert list(result["subtype_kmeans"]) == [f"{s}_0" for s in obs["subtype"]] + + def test_refinement_produces_budget_many_distinct_sublabels(self): + # 3 clearly-separable blocks so K-means should cleanly find 3 + # clusters, none dropped as noise. + block1 = _make_expression(100, 20, seed=1) + block2 = _make_expression(100, 20, seed=2) + block2[:, :5] += 30 + block3 = _make_expression(100, 20, seed=3) + block3[:, 5:10] += 30 + X = np.vstack([block1, block2, block3]) + obs = pd.DataFrame({"major": ["A"] * 300, "subtype": ["A_0"] * 300}) + plan = pd.DataFrame({ + "major": ["A"], "subtype": ["A_0"], "n_cells": [300], + "current_k": [1], "expected_k": [3], "budget": [3], + }) + result = refine_predefined_subtypes(X, obs, ["major", "subtype"], plan, min_frac=0.0, random_state=0) + assert set(result["subtype_kmeans"]) == {"A_0_0", "A_0_1", "A_0_2"} + + def test_gene_space_shared_across_labels_in_same_major_type(self, capsys): + # A has its own marker block (vs B) so real narrowing occurs, not + # the "0 others" fallback that would trivially return all genes. + X_a = _make_expression(200, 600, seed=4) + X_a[:, :50] += 20 + X_b = _make_expression(200, 600, seed=5) + X = np.vstack([X_a, X_b]) + obs = pd.DataFrame({ + "major": ["A"] * 200 + ["B"] * 200, + "subtype": ["A_0"] * 100 + ["A_1"] * 100 + ["B_0"] * 200, + }) + plan = pd.DataFrame({ + "major": ["A", "A", "B"], "subtype": ["A_0", "A_1", "B_0"], + "n_cells": [100, 100, 200], "current_k": [2, 2, 1], "expected_k": [2, 2, 1], + "budget": [2, 2, 1], + }) + refine_predefined_subtypes(X, obs, ["major", "subtype"], plan, min_frac=0.0, random_state=0, verbose=True, max_genes = 500) + out = capsys.readouterr().out + counts = re.findall(r"using (\d+) genes", out) + assert len(counts) == 2 + assert counts[0] == counts[1] + assert counts[0] != "600" + + def test_finest_label_sharing_name_with_own_major_type_does_not_collide(self, capsys): + # ATL's own finest label is named identically to its major type -- + # a real case from the user's atlas. The gene-narrowing comparator + # dict is keyed purely by major-type value (never by finest + # label), so this must resolve exactly like any other label name. + X_atl = _make_expression(200, 600, seed=44) + X_atl[:, :50] += 20 + X_pt = _make_expression(200, 600, seed=45) + X = np.vstack([X_atl, X_pt]) + obs = pd.DataFrame({ + "major": ["ATL"] * 200 + ["PT"] * 200, + "subtype": ["ATL"] * 200 + ["PT_0"] * 200, + }) + plan = pd.DataFrame({ + "major": ["ATL", "PT"], "subtype": ["ATL", "PT_0"], + "n_cells": [200, 200], "current_k": [1, 1], "expected_k": [2, 1], + "budget": [2, 1], + }) + result = refine_predefined_subtypes( + X, obs, ["major", "subtype"], plan, min_frac=0.0, random_state=0, verbose=True, + ) + out = capsys.readouterr().out + assert "using" in out + assert set(result["subtype_kmeans"].iloc[:200]) == {"ATL_0", "ATL_1"} + assert set(result["subtype_kmeans"].iloc[200:]) == {"PT_0_0"} + + def test_partial_noise_drop_leaves_dropped_cells_none(self): + # 5 cells with a clearly shifted profile among 295 background + # cells -- forced K=2 makes the shifted group its own cluster, + # which should then be dropped as noise. + background = _make_expression(295, 20, seed=12) + rare = _make_expression(5, 20, seed=13) + rare[:, :5] += 50 + X = np.vstack([background, rare]) + obs = pd.DataFrame({"major": ["A"] * 300, "subtype": ["A_0"] * 300}) + plan = pd.DataFrame({ + "major": ["A"], "subtype": ["A_0"], "n_cells": [300], + "current_k": [1], "expected_k": [2], "budget": [2], + }) + with pytest.warns(UserWarning, match=r"Dropped 5 cell"): + result = refine_predefined_subtypes(X, obs, ["major", "subtype"], plan, min_frac=0.025, random_state=0) + assert result["subtype_kmeans"].iloc[:295].notna().all() + assert result["subtype_kmeans"].iloc[295:].isna().all() + + def test_all_dropped_fallback_covers_cells_excluded_by_cap(self): + X = _make_expression(300, 20, seed=7) + obs = pd.DataFrame({"major": ["A"] * 300, "subtype": ["A_0"] * 300}) + plan = pd.DataFrame({ + "major": ["A"], "subtype": ["A_0"], "n_cells": [300], + "current_k": [1], "expected_k": [5], "budget": [5], + }) + with pytest.warns(UserWarning, match=r"All candidate subtypes"): + result = refine_predefined_subtypes( + X, obs, ["major", "subtype"], plan, + min_frac=0.99, max_cells_per_label=100, random_state=0, + ) + # All 300 cells get the fallback label, not just the 100 sampled. + assert result["subtype_kmeans"].notna().all() + assert set(result["subtype_kmeans"]) == {"A_0_0"} + + def test_min_cells_floor_drops_clusters_below_ten(self): + # min_frac=0.001 makes the fraction-based threshold negligible + # (<1 cell); only the absolute floor of 10 should decide outcomes. + background = _make_expression(270, 20, seed=30) + group15 = _make_expression(15, 20, seed=31) + group15[:, :5] += 40 + group5 = _make_expression(5, 20, seed=32) + group5[:, 5:10] += 40 + X = np.vstack([background, group15, group5]) + obs = pd.DataFrame({"major": ["A"] * 290, "subtype": ["A_0"] * 290}) + plan = pd.DataFrame({ + "major": ["A"], "subtype": ["A_0"], "n_cells": [290], + "current_k": [1], "expected_k": [3], "budget": [3], + }) + with pytest.warns(UserWarning, match=r"Dropped 5 cell"): + result = refine_predefined_subtypes(X, obs, ["major", "subtype"], plan, min_frac=0.001, random_state=0) + assert result["subtype_kmeans"].iloc[270:285].notna().all() # 15-cell group survives (>=10) + assert result["subtype_kmeans"].iloc[285:290].isna().all() # 5-cell group dropped (<10) + + def test_max_cells_per_label_caps_successful_clustering(self): + block_a = _make_expression(300, 20, seed=20) + block_b = _make_expression(300, 20, seed=21) + block_b[:, :5] += 30 + X = np.vstack([block_a, block_b]) + obs = pd.DataFrame({"major": ["A"] * 600, "subtype": ["A_0"] * 600}) + plan = pd.DataFrame({ + "major": ["A"], "subtype": ["A_0"], "n_cells": [600], + "current_k": [1], "expected_k": [2], "budget": [2], + }) + result = refine_predefined_subtypes( + X, obs, ["major", "subtype"], plan, + min_frac=0.0, max_cells_per_label=100, random_state=0, + ) + assert result["subtype_kmeans"].notna().sum() == 100 + + def test_hierarchy_consistency_holds_for_mixed_flagged_and_unflagged(self): + obs_a = _make_obs({"A": [11284, 3141]}) # under-clustered, gets refined + obs_a["major"] = "A" + obs_b = _make_obs({"B": [500] * 14}) # already well-resolved (in fact over-clustered), untouched + obs_b["major"] = "B" + obs = pd.concat([obs_a, obs_b], ignore_index=True) + X = _make_expression(len(obs), 20, seed=40) + # A under-clustered and B over-clustered are both intrinsic to this + # mixed scenario, so both warnings are expected here, not incidental. + with pytest.warns(UserWarning) as record: + plan = plan_subtype_refinement(obs, ["major", "subtype"], return_counts=True) + assert len(record) == 2 + result = refine_predefined_subtypes(X, obs, ["major", "subtype"], plan, min_frac=0.0, random_state=0) + assigned = result.dropna(subset=["subtype_kmeans"]) + check_subtype_consistency(assigned, ["major", "subtype", "subtype_kmeans"]) # no raise + + def test_gene_narrowing_skipped_when_n_genes_within_max_genes(self, capsys): + X = _make_expression(300, 20, seed=50) + obs = pd.DataFrame({"major": ["A"] * 300, "subtype": ["A_0"] * 300}) + plan = pd.DataFrame({ + "major": ["A"], "subtype": ["A_0"], "n_cells": [300], + "current_k": [1], "expected_k": [2], "budget": [2], + }) + refine_predefined_subtypes(X, obs, ["major", "subtype"], plan, min_frac=0.0, random_state=0, verbose=True, max_genes = 500) + out = capsys.readouterr().out + assert "using 20 genes" in out + + def test_output_shape_and_index(self): + X = _make_expression(100, 10, seed=60) + obs = pd.DataFrame({"major": ["A"] * 50 + ["B"] * 50, "subtype": ["A_0"] * 50 + ["B_0"] * 50}) + plan = pd.DataFrame({ + "major": ["A", "B"], "subtype": ["A_0", "B_0"], "n_cells": [50, 50], + "current_k": [1, 1], "expected_k": [1, 1], "budget": [1, 1], + }) + result = refine_predefined_subtypes(X, obs, ["major", "subtype"], plan) + assert list(result.columns) == ["major", "subtype", "subtype_kmeans"] + assert result.shape[0] == 100 + assert list(result.index) == list(obs.index) + + def test_reproducible_with_fixed_random_state(self): + X = _make_expression(300, 20, seed=70) + obs = pd.DataFrame({"major": ["A"] * 300, "subtype": ["A_0"] * 300}) + plan = pd.DataFrame({ + "major": ["A"], "subtype": ["A_0"], "n_cells": [300], + "current_k": [1], "expected_k": [3], "budget": [3], + }) + r1 = refine_predefined_subtypes(X, obs, ["major", "subtype"], plan, min_frac=0.0, random_state=5) + r2 = refine_predefined_subtypes(X, obs, ["major", "subtype"], plan, min_frac=0.0, random_state=5) + assert (r1["subtype_kmeans"].fillna("NA") == r2["subtype_kmeans"].fillna("NA")).all() + + def test_sparse_and_dense_input_agree(self): + X_dense = _make_expression(300, 20, seed=80) + X_sparse = csr_matrix(X_dense) + obs = pd.DataFrame({"major": ["A"] * 300, "subtype": ["A_0"] * 300}) + plan = pd.DataFrame({ + "major": ["A"], "subtype": ["A_0"], "n_cells": [300], + "current_k": [1], "expected_k": [2], "budget": [2], + }) + r_dense = refine_predefined_subtypes(X_dense, obs, ["major", "subtype"], plan, min_frac=0.0, random_state=0) + r_sparse = refine_predefined_subtypes(X_sparse, obs, ["major", "subtype"], plan, min_frac=0.0, random_state=0) + assert (r_dense["subtype_kmeans"].fillna("NA") == r_sparse["subtype_kmeans"].fillna("NA")).all() + + +class TestSummarizeSubtypes: + def test_basic_output_shape_and_ratios(self): + labels = _make_obs({"A": [10, 10], "B": [10]}) + X = _make_expression(30, 5, seed=0) + summary = summarize_subtypes(X, labels) + assert summary["centroids"].shape == (3, 5) + assert list(summary["hierarchy"].columns) == ["major", "subtype"] + assert len(summary["hierarchy"]) == 3 + assert set(summary["ratios"].keys()) == {"A", "B"} + assert summary["ratios"]["A"] == pytest.approx(20 / 30) + assert summary["ratios"]["B"] == pytest.approx(10 / 30) + assert sum(summary["ratios"].values()) == pytest.approx(1.0) + + def test_centroid_values_are_correct_means(self): + labels = pd.DataFrame({ + "major": ["A"] * 6, + "subtype": ["A_0"] * 3 + ["A_1"] * 3, + }) + X = np.array([ + [1.0, 2.0], + [3.0, 4.0], + [5.0, 6.0], + [10.0, 20.0], + [30.0, 40.0], + [50.0, 60.0], + ]) + summary = summarize_subtypes(X, labels) + subtype_col = summary["hierarchy"]["subtype"].to_numpy() + idx_a0 = int(np.where(subtype_col == "A_0")[0][0]) + idx_a1 = int(np.where(subtype_col == "A_1")[0][0]) + assert np.allclose(summary["centroids"][idx_a0], [3.0, 4.0]) + assert np.allclose(summary["centroids"][idx_a1], [30.0, 40.0]) + + def test_hierarchy_crosswalk_matches_each_groups_ancestry(self): + labels = pd.DataFrame({ + "major": ["T"] * 12, + "meso": ["T_m0"] * 6 + ["T_m1"] * 6, + "fine": ["T_m0_f0"] * 3 + ["T_m0_f1"] * 3 + ["T_m1_f0"] * 3 + ["T_m1_f1"] * 3, + }) + X = _make_expression(12, 4, seed=1) + summary = summarize_subtypes(X, labels) + hierarchy = summary["hierarchy"] + assert len(hierarchy) == 4 + + expected_ancestry = { + "T_m0_f0": ("T", "T_m0"), + "T_m0_f1": ("T", "T_m0"), + "T_m1_f0": ("T", "T_m1"), + "T_m1_f1": ("T", "T_m1"), + } + for _, row in hierarchy.iterrows(): + assert (row["major"], row["meso"]) == expected_ancestry[row["fine"]] + + # centroids stay row-aligned with hierarchy + idx_f0 = int(np.where(hierarchy["fine"].to_numpy() == "T_m0_f0")[0][0]) + assert np.allclose(summary["centroids"][idx_f0], X[:3].mean(axis=0)) + + def test_four_level_hierarchy_supported(self): + labels = pd.DataFrame({ + "l0": ["X"] * 8, + "l1": ["X_a"] * 4 + ["X_b"] * 4, + "l2": ["X_a_i"] * 4 + ["X_b_i"] * 4, + "l3": ["X_a_i_1"] * 4 + ["X_b_i_1"] * 4, + }) + X = _make_expression(8, 3, seed=2) + summary = summarize_subtypes(X, labels) + assert list(summary["hierarchy"].columns) == ["l0", "l1", "l2", "l3"] + assert summary["centroids"].shape == (2, 3) + assert set(summary["ratios"].keys()) == {"X"} + + def test_missing_finest_label_excluded(self): + labels = pd.DataFrame({ + "major": ["A"] * 10, + "subtype": ["A_0"] * 7 + [None] * 3, + }) + X = _make_expression(10, 4, seed=3) + summary = summarize_subtypes(X, labels) + assert summary["centroids"].shape == (1, 4) + assert np.allclose(summary["centroids"][0], X[:7].mean(axis=0)) + assert list(summary["hierarchy"]["subtype"]) == ["A_0"] + + def test_ratios_use_full_original_coarsest_count(self): + # major "A" has 10 cells total, only 7 keep a finest-level label; + # its ratio must reflect 10, not 7. + labels = pd.DataFrame({ + "major": ["A"] * 10 + ["B"] * 5, + "subtype": ["A_0"] * 7 + [None] * 3 + ["B_0"] * 5, + }) + X = _make_expression(15, 4, seed=4) + summary = summarize_subtypes(X, labels) + assert summary["ratios"]["A"] == pytest.approx(10 / 15) + assert summary["ratios"]["B"] == pytest.approx(5 / 15) + + def test_fully_dropped_coarsest_group_excluded_and_renormalized(self): + labels = pd.DataFrame({ + "major": ["A"] * 5 + ["B"] * 5, + "subtype": [None] * 5 + ["B_0"] * 5, + }) + X = _make_expression(10, 4, seed=5) + summary = summarize_subtypes(X, labels) + assert set(summary["ratios"].keys()) == {"B"} + assert summary["ratios"]["B"] == pytest.approx(1.0) + assert summary["centroids"].shape == (1, 4) + + @pytest.mark.parametrize("n_cells", [1, 2]) + def test_small_group_warns_below_three_cells(self, n_cells): + labels = pd.DataFrame({ + "major": ["A"] * n_cells, + "subtype": ["A_0"] * n_cells, + }) + X = _make_expression(n_cells, 4, seed=6) + with pytest.warns(UserWarning, match=r"has only \d+ cell"): + summarize_subtypes(X, labels) + + def test_no_warning_at_exactly_three_cells(self, recwarn): + labels = pd.DataFrame({ + "major": ["A"] * 3, + "subtype": ["A_0"] * 3, + }) + X = _make_expression(3, 4, seed=7) + summarize_subtypes(X, labels) + assert len(recwarn) == 0 + + def test_inconsistent_hierarchy_raises(self): + labels = pd.DataFrame({ + "major": ["A", "B"], + "subtype": ["S1", "S1"], + }) + X = _make_expression(2, 4, seed=8) + with pytest.raises(ValueError, match="Inconsistent hierarchy"): + summarize_subtypes(X, labels) + + def test_sparse_and_dense_agree(self): + labels = _make_obs({"A": [5, 5], "B": [5]}) + X_dense = _make_expression(15, 4, seed=9) + X_sparse = csr_matrix(X_dense) + + summary_dense = summarize_subtypes(X_dense, labels) + summary_sparse = summarize_subtypes(X_sparse, labels) + + assert isinstance(summary_dense["centroids"], np.ndarray) + assert issparse(summary_sparse["centroids"]) + assert np.allclose(summary_dense["centroids"], summary_sparse["centroids"].toarray()) + + def test_verbose_prints_summary(self, capsys): + labels = _make_obs({"A": [5, 5]}) + X = _make_expression(10, 4, seed=10) + summarize_subtypes(X, labels, verbose=True) + out = capsys.readouterr().out + assert "10 cells" in out + assert "2 'subtype' groups" in out + assert "1 'major' groups" in out + + def test_verbose_false_prints_nothing(self, capsys): + labels = _make_obs({"A": [5, 5]}) + X = _make_expression(10, 4, seed=11) + summarize_subtypes(X, labels, verbose=False) + out = capsys.readouterr().out + assert out == "" + + def test_composes_with_kmeans_define_subtypes_output(self): + X = _make_expression(300, 20, seed=12) + annotations = np.array(["A"] * 150 + ["B"] * 150) + labels = kmeans_define_subtypes(X, annotations, subcluster_size=3, random_state=0) + summary = summarize_subtypes(X, labels) + assert summary["centroids"].shape[1] == 20 + assert set(summary["hierarchy"]["major"].unique()) <= {"A", "B"} + assert sum(summary["ratios"].values()) == pytest.approx(1.0) + + def test_composes_with_predefined_subtypes_output(self): + # repeat the fixture 4x so every fine-level group has 4 cells, + # comfortably above the <3-cell warning threshold. Each major type + # now has 16 cells and 4 distinct fine labels, matched exactly by + # k_heuristic to avoid an incidental under-clustering warning. + obs_wide = pd.concat([_consistent_3level_obs()] * 4, ignore_index=True) + X = _make_expression(len(obs_wide), 6, seed=13) + labels = predefined_subtypes( + obs_wide, ["major", "meso", "fine"], min_frac=0.0, k_heuristic=lambda n: 4 + ) + summary = summarize_subtypes(X, labels) + assert list(summary["hierarchy"].columns) == ["major", "meso", "fine"] + assert summary["centroids"].shape[1] == 6 + assert sum(summary["ratios"].values()) == pytest.approx(1.0) + + +class TestDeScoresLogfcRank: + def test_output_shape_and_dtype(self): + centroids = _make_expression(3, 10, seed=0) + scores = _de_scores_logfc_rank(centroids) + assert scores.shape == (3, 10) + assert scores.dtype == np.float32 + + def test_matches_hand_verified_values_no_ties(self): + # Every gene has a distinct log-fold-change within each row, so the + # ranking is fully deterministic (no tie-breaking involved). + centroids = np.array([ + [20.0, 1.0, 8.0, 4.0], + [1.0, 20.0, 4.0, 8.0], + ]) + scores = _de_scores_logfc_rank(centroids) + assert scores.tolist() == [[1.0, 4.0, 2.0, 3.0], [4.0, 1.0, 3.0, 2.0]] + + def test_lower_score_means_more_distinctive_marker(self): + # gene0 is cluster0's marker, gene1 is cluster1's -- each cluster's + # own marker gene should get the lowest (best) score in its row. + centroids = np.array([ + [10.0, 1.0, 5.0, 5.0], + [1.0, 10.0, 5.0, 5.0], + ]) + scores = _de_scores_logfc_rank(centroids) + assert scores[0].argmin() == 0 + assert scores[1].argmin() == 1 + + def test_median_over_more_than_one_comparison(self): + # 3 clusters -> each row's score is a median over 2 pairwise + # comparisons, not a single one. + centroids = np.array([ + [10.0, 1.0, 1.0], + [1.0, 10.0, 1.0], + [1.0, 1.0, 10.0], + ]) + scores = _de_scores_logfc_rank(centroids) + assert scores.argmin(axis=1).tolist() == [0, 1, 2] + + def test_zero_entries_do_not_produce_nan_or_inf(self): + centroids = np.array([[0.0, 5.0], [5.0, 0.0]]) + scores = _de_scores_logfc_rank(centroids) + assert not np.isnan(scores).any() + assert not np.isinf(scores).any() + + +class TestSelectDeGenes: + def test_selects_genes_with_known_differential_signal(self): + centroids = _make_expression(3, 20, seed=0) + centroids[0, :3] += 1000 + genes = select_de_genes(centroids, max_genes=3) + assert set(genes.tolist()) == {0, 1, 2} + + def test_passthrough_when_n_genes_leq_max_genes(self): + centroids = _make_expression(3, 5, seed=1) + genes = select_de_genes(centroids, max_genes=10) + assert list(genes) == list(range(5)) + + def test_passthrough_return_scores_is_none(self): + centroids = _make_expression(3, 5, seed=2) + genes, scores = select_de_genes(centroids, max_genes=10, return_scores=True) + assert list(genes) == list(range(5)) + assert scores is None + + def test_return_scores_matches_selected_genes(self): + centroids = _make_expression(4, 30, seed=3) + genes, scores = select_de_genes(centroids, max_genes=8, return_scores=True) + assert scores.shape == (4, 30) + expected_genes = np.argsort(scores.min(axis=0))[:8] + assert np.array_equal(genes, expected_genes) + + def test_custom_callable_is_used(self): + centroids = _make_expression(3, 20, seed=4) + + def force_gene_5_most_distinctive(c): + scores = np.ones_like(c) + scores[:, 5] = 0.0 + return scores + + genes = select_de_genes(centroids, max_genes=3, method=force_gene_5_most_distinctive) + assert 5 in genes + + def test_custom_callable_receives_dense_array(self): + centroids_dense = _make_expression(3, 20, seed=5) + centroids_sparse = csr_matrix(centroids_dense) + received = {} + + def probe(c): + received["is_sparse"] = issparse(c) + received["shape"] = c.shape + return _make_expression(*c.shape, seed=6) + + select_de_genes(centroids_sparse, max_genes=3, method=probe) + assert received["is_sparse"] is False + assert received["shape"] == (3, 20) + + def test_custom_callable_not_called_in_passthrough(self): + centroids = _make_expression(3, 5, seed=7) + calls = [] + + def probe(c): + calls.append(c) + return np.zeros_like(c) + + select_de_genes(centroids, max_genes=10, method=probe) + assert len(calls) == 0 + + def test_unknown_method_raises(self): + centroids = _make_expression(3, 20, seed=8) + with pytest.raises(ValueError, match="Unknown method"): + select_de_genes(centroids, max_genes=3, method="bogus") + + def test_sparse_and_dense_agree(self): + centroids_dense = _make_expression(4, 30, seed=9) + centroids_sparse = csr_matrix(centroids_dense) + genes_dense = select_de_genes(centroids_dense, max_genes=8) + genes_sparse = select_de_genes(centroids_sparse, max_genes=8) + assert np.array_equal(genes_dense, genes_sparse) + + def test_verbose_prints_summary(self, capsys): + centroids = _make_expression(3, 20, seed=10) + select_de_genes(centroids, max_genes=5, verbose=True) + out = capsys.readouterr().out + assert "Computing gene scores" in out + assert "Selected 5 DE genes" in out + + def test_verbose_false_prints_nothing(self, capsys): + centroids = _make_expression(3, 20, seed=11) + select_de_genes(centroids, max_genes=5, verbose=False) + out = capsys.readouterr().out + assert out == "" + + +def _valid_ref_adata(n_cells=20, n_genes=10, seed=0): + obs = pd.DataFrame({"cell_type": ["T0"] * (n_cells // 2) + ["T1"] * (n_cells // 2)}) + return _make_ref_adata(obs, n_genes=n_genes, seed=seed) + + +class TestValidateReferenceInput: + def test_valid_input_does_not_raise(self): + adata = _valid_ref_adata() + validate_reference_input( + adata, "cell_type", max_input_genes=100, + warn_mt=True, check_counts=True, check_nonnegative=True, + ) # no raise + + def test_missing_cell_type_key_raises(self): + adata = _valid_ref_adata() + with pytest.raises(ValueError, match="not found in adata.obs"): + validate_reference_input( + adata, "bogus_key", max_input_genes=100, + warn_mt=True, check_counts=True, check_nonnegative=True, + ) + + def test_nan_in_cell_type_key_raises(self): + adata = _valid_ref_adata() + adata.obs.loc[adata.obs_names[0], "cell_type"] = np.nan + with pytest.raises(ValueError, match="missing \\(NaN\\) value"): + validate_reference_input( + adata, "cell_type", max_input_genes=100, + warn_mt=True, check_counts=True, check_nonnegative=True, + ) + + def test_list_of_keys_checks_every_key_not_just_the_first(self): + adata = _valid_ref_adata() + adata.obs["lineage"] = "L0" + adata.obs.loc[adata.obs_names[0], "lineage"] = np.nan + with pytest.raises(ValueError, match="lineage.*missing \\(NaN\\) value"): + validate_reference_input( + adata, ["cell_type", "lineage"], max_input_genes=100, + warn_mt=True, check_counts=True, check_nonnegative=True, + ) + + def test_negative_values_raise_by_default(self): + adata = _valid_ref_adata() + adata.X[0, 0] = -1.0 + with pytest.raises(ValueError, match="negative values"): + validate_reference_input( + adata, "cell_type", max_input_genes=100, + warn_mt=True, check_counts=True, check_nonnegative=True, + ) + + def test_check_nonnegative_false_bypasses_negative_check(self): + adata = _valid_ref_adata() + adata.X[0, 0] = -1.0 + validate_reference_input( + adata, "cell_type", max_input_genes=100, + warn_mt=True, check_counts=True, check_nonnegative=False, + ) # no raise + + def test_non_integer_counts_raise_by_default(self): + adata = _valid_ref_adata() + adata.X[0, 0] = 1.5 + with pytest.raises(ValueError, match="does not look like raw counts"): + validate_reference_input( + adata, "cell_type", max_input_genes=100, + warn_mt=True, check_counts=True, check_nonnegative=True, + ) + + def test_check_counts_false_bypasses_noninteger_check(self): + adata = _valid_ref_adata() + adata.X[0, 0] = 1.5 + validate_reference_input( + adata, "cell_type", max_input_genes=100, + warn_mt=True, check_counts=False, check_nonnegative=True, + ) # no raise + + def test_empty_cell_raises(self): + adata = _valid_ref_adata() + adata.X[0, :] = 0 + with pytest.raises(ValueError, match="cell\\(s\\) have zero total counts"): + validate_reference_input( + adata, "cell_type", max_input_genes=100, + warn_mt=True, check_counts=True, check_nonnegative=True, + ) + + def test_empty_gene_raises(self): + adata = _valid_ref_adata() + adata.X[:, 0] = 0 + with pytest.raises(ValueError, match="gene\\(s\\) have zero total counts"): + validate_reference_input( + adata, "cell_type", max_input_genes=100, + warn_mt=True, check_counts=True, check_nonnegative=True, + ) + + def test_warn_mt_true_warns_on_mt_genes(self): + adata = _valid_ref_adata() + adata.var_names = ["MT-ND1"] + list(adata.var_names[1:]) + with pytest.warns(UserWarning, match="MT-"): + validate_reference_input( + adata, "cell_type", max_input_genes=100, + warn_mt=True, check_counts=True, check_nonnegative=True, + ) + + def test_warn_mt_false_suppresses_warning(self, recwarn): + adata = _valid_ref_adata() + adata.var_names = ["MT-ND1"] + list(adata.var_names[1:]) + validate_reference_input( + adata, "cell_type", max_input_genes=100, + warn_mt=False, check_counts=True, check_nonnegative=True, + ) + assert len(recwarn) == 0 + + def test_too_many_genes_raises(self): + adata = _valid_ref_adata(n_genes=10) + with pytest.raises(ValueError, match="exceeding max_input_genes"): + validate_reference_input( + adata, "cell_type", max_input_genes=5, + warn_mt=True, check_counts=True, check_nonnegative=True, + ) + + +class TestSetupReference: + def test_kmeans_path_output_structure(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 50 + ["T1"] * 50}) + adata = _make_ref_adata(obs, n_genes=20, seed=0) + result = setup_reference( + adata, cell_type_key="cell_type", subcluster_size=3, random_state=0, max_input_genes=100 + ) + assert isinstance(result, AnnData) + assert list(result.obs.columns) == ["cell_type", "subtype"] + assert list(result.obs_names) == list(result.obs["subtype"]) + assert set(result.uns["ratios"].keys()) == {"T0", "T1"} + assert sum(result.uns["ratios"].values()) == pytest.approx(1.0) + assert list(result.uns["level_keys"]) == ["cell_type", "subtype"] + + def test_string_and_single_element_list_are_equivalent(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 50 + ["T1"] * 50}) + adata = _make_ref_adata(obs, n_genes=20, seed=1) + r_str = setup_reference( + adata, cell_type_key="cell_type", subcluster_size=3, random_state=0, max_input_genes=100 + ) + r_list = setup_reference( + adata, cell_type_key=["cell_type"], subcluster_size=3, random_state=0, max_input_genes=100 + ) + assert list(r_str.obs_names) == list(r_list.obs_names) + assert np.allclose(r_str.X, r_list.X) + assert r_str.uns["ratios"] == r_list.uns["ratios"] + + def test_predefined_path_output_structure(self, recwarn): + # kws_predefined overrides k_heuristic to match the 2 labels each + # major type actually has, deterministically clearing both the + # over- and under-clustering warnings regardless of cell counts. + obs = _make_obs({"T0": [25, 25], "T1": [25, 25]}) + adata = _make_ref_adata(obs, n_genes=20, seed=2) + result = setup_reference( + adata, cell_type_key=["major", "subtype"], max_input_genes=100, min_frac=0.0, + kws_predefined={"k_heuristic": lambda n: 2}, + ) + assert list(result.obs.columns) == ["major", "subtype"] + assert set(result.obs["subtype"]) == {"T0_0", "T0_1", "T1_0", "T1_1"} + assert list(result.obs_names) == list(result.obs["subtype"]) + assert list(result.uns["level_keys"]) == ["major", "subtype"] + assert len(recwarn) == 0 + + def test_predefined_path_three_levels(self, recwarn): + n_per_type = 40 + meso = sum([[f"T{i}_m{j % 2}" for j in range(n_per_type)] for i in range(2)], []) + fine = sum([[f"T{i}_m{j % 2}_f{j % 4}" for j in range(n_per_type)] for i in range(2)], []) + obs = pd.DataFrame({ + "cell_type": ["T0"] * n_per_type + ["T1"] * n_per_type, + "meso": meso, + "fine": fine, + }) + adata = _make_ref_adata(obs, n_genes=20, seed=3) + # k_heuristic overridden to match the 4 distinct fine labels each + # cell_type actually has, clearing both the over- and + # under-clustering warnings deterministically. + result = setup_reference( + adata, cell_type_key=["cell_type", "meso", "fine"], max_input_genes=100, min_frac=0.0, + kws_predefined={"k_heuristic": lambda n: 4}, + ) + assert list(result.obs.columns) == ["cell_type", "meso", "fine"] + assert list(result.obs_names) == list(result.obs["fine"]) + assert list(result.uns["level_keys"]) == ["cell_type", "meso", "fine"] + assert len(recwarn) == 0 + + def test_max_genes_limits_output_gene_count(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 50 + ["T1"] * 50}) + adata = _make_ref_adata(obs, n_genes=30, seed=4) + result = setup_reference( + adata, cell_type_key="cell_type", subcluster_size=3, max_genes=10, + max_input_genes=100, random_state=0, + ) + assert result.shape[1] == 10 + assert len(result.var_names) == 10 + + def test_max_input_genes_raises_when_exceeded(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 50 + ["T1"] * 50}) + adata = _make_ref_adata(obs, n_genes=30, seed=5) + with pytest.raises(ValueError, match="max_input_genes"): + setup_reference(adata, cell_type_key="cell_type", max_input_genes=10) + + def test_missing_coarsest_column_raises(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 10 + ["T1"] * 10}) + adata = _make_ref_adata(obs, n_genes=10, seed=6) + with pytest.raises(ValueError, match="not found in adata.obs"): + setup_reference(adata, cell_type_key="bogus_col", max_input_genes=100) + + def test_missing_finer_column_raises(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 10 + ["T1"] * 10}) + adata = _make_ref_adata(obs, n_genes=10, seed=7) + with pytest.raises(ValueError, match="not found in adata.obs"): + setup_reference(adata, cell_type_key=["cell_type", "bogus_subtype"], max_input_genes=100) + + def test_nan_in_cell_type_column_raises(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 10 + ["T1"] * 10}) + adata = _make_ref_adata(obs, n_genes=10, seed=8) + adata.obs.loc[adata.obs.index[0], "cell_type"] = None + with pytest.raises(ValueError, match=r"missing \(NaN\)"): + setup_reference(adata, cell_type_key="cell_type", max_input_genes=100) + + def test_ratios_match_expected_fractions(self): + obs = pd.DataFrame({"cell_type": ["A"] * 75 + ["B"] * 25}) + adata = _make_ref_adata(obs, n_genes=10, seed=9) + result = setup_reference( + adata, cell_type_key="cell_type", subcluster_size=3, max_input_genes=100, random_state=0 + ) + assert result.uns["ratios"]["A"] == pytest.approx(0.75) + assert result.uns["ratios"]["B"] == pytest.approx(0.25) + + def test_sparse_X_works_end_to_end(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 50 + ["T1"] * 50}) + adata_dense = _make_ref_adata(obs, n_genes=20, seed=10) + adata_sparse = adata_dense.copy() + adata_sparse.X = csr_matrix(adata_dense.X) + r_dense = setup_reference( + adata_dense, cell_type_key="cell_type", subcluster_size=3, random_state=0, max_input_genes=100 + ) + r_sparse = setup_reference( + adata_sparse, cell_type_key="cell_type", subcluster_size=3, random_state=0, max_input_genes=100 + ) + X_sparse_dense = r_sparse.X.toarray() if issparse(r_sparse.X) else r_sparse.X + assert np.allclose(r_dense.X, X_sparse_dense) + + def test_verbose_prints_progress(self, capsys): + obs = pd.DataFrame({"cell_type": ["T0"] * 50 + ["T1"] * 50}) + adata = _make_ref_adata(obs, n_genes=20, seed=11) + setup_reference( + adata, cell_type_key="cell_type", subcluster_size=3, random_state=0, + max_input_genes=100, verbose=True, + ) + out = capsys.readouterr().out + assert "Defining subtypes" in out + assert "Reference prepared" in out + + def test_verbose_false_prints_nothing(self, capsys): + obs = pd.DataFrame({"cell_type": ["T0"] * 50 + ["T1"] * 50}) + adata = _make_ref_adata(obs, n_genes=20, seed=12) + setup_reference( + adata, cell_type_key="cell_type", subcluster_size=3, random_state=0, + max_input_genes=100, verbose=False, + ) + out = capsys.readouterr().out + assert out == "" + + def test_h5ad_round_trip_preserves_everything(self, tmp_path): + obs = pd.DataFrame({"cell_type": ["T0"] * 50 + ["T1"] * 50}) + adata = _make_ref_adata(obs, n_genes=20, seed=13) + result = setup_reference( + adata, cell_type_key="cell_type", subcluster_size=3, random_state=0, max_input_genes=100 + ) + path = tmp_path / "ref.h5ad" + result.write_h5ad(path) + loaded = sc.read_h5ad(path) + + assert list(loaded.obs.columns) == list(result.obs.columns) + assert list(loaded.obs_names) == list(result.obs_names) + assert list(loaded.var_names) == list(result.var_names) + assert np.allclose(np.asarray(loaded.X), np.asarray(result.X)) + assert loaded.uns["ratios"] == result.uns["ratios"] + assert list(loaded.uns["level_keys"]) == list(result.uns["level_keys"]) + + def test_subcluster_size_affects_output_subtype_count(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 300}) + adata = _make_ref_adata(obs, n_genes=20, seed=14) + result_capped = setup_reference( + adata, cell_type_key="cell_type", subcluster_size=2, random_state=0, max_input_genes=100 + ) + result_loose = setup_reference( + adata, cell_type_key="cell_type", subcluster_size=10, random_state=0, max_input_genes=100 + ) + assert result_capped.shape[0] <= 2 + assert result_capped.shape[0] < result_loose.shape[0] + + def test_kws_kmeans_forwards_to_kmeans_define_subtypes(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 300}) + adata = _make_ref_adata(obs, n_genes=20, seed=15) + result = setup_reference( + adata, cell_type_key="cell_type", subcluster_size=10, random_state=0, max_input_genes=100, + kws_kmeans={"max_cells_per_type": 50}, + ) + # only 50 of the 300 cells get sampled/assigned, so no subtype can + # have accumulated more than 50 cells' worth of representation. + assert result.shape[0] >= 1 + + def test_kws_predefined_k_heuristic_silences_overclustering_warning(self, recwarn): + # Default k_heuristic on 20 cells gives expected_k=1, so the 4 + # labels here would normally overcluster; overriding to match the + # 4 labels exactly proves the override is wired through (and, + # since the check is symmetric, also avoids under-clustering). + obs = _make_obs({"T0": [5, 5, 5, 5]}) + adata = _make_ref_adata(obs, n_genes=20, seed=16) + setup_reference( + adata, cell_type_key=["major", "subtype"], max_input_genes=100, min_frac=0.0, + kws_predefined={"k_heuristic": lambda n: 4}, + ) + assert len(recwarn) == 0 + + def test_kws_de_genes_return_scores_stored_in_uns(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 50 + ["T1"] * 50}) + adata = _make_ref_adata(obs, n_genes=20, seed=17) + result = setup_reference( + adata, cell_type_key="cell_type", subcluster_size=3, max_genes=5, + random_state=0, max_input_genes=100, kws_de_genes={"return_scores": True}, + ) + assert "de_scores" in result.uns + assert result.uns["de_scores"].shape == (result.shape[0], 20) + + def test_kws_duplicate_parameter_raises(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 50 + ["T1"] * 50}) + adata = _make_ref_adata(obs, n_genes=20, seed=18) + with pytest.raises(TypeError, match="subcluster_size"): + setup_reference( + adata, cell_type_key="cell_type", subcluster_size=3, max_input_genes=100, + kws_kmeans={"subcluster_size": 5}, + ) + + def test_predefined_path_without_refinement_has_subtype_plan_but_no_n_subtypes_made(self): + obs = _make_obs({"T0": [25, 25], "T1": [25, 25]}) + adata = _make_ref_adata(obs, n_genes=20, seed=19) + result = setup_reference( + adata, cell_type_key=["major", "subtype"], max_input_genes=100, min_frac=0.0, + kws_predefined={"k_heuristic": lambda n: 2}, + ) + assert "subtype_plan" in result.uns + assert list(result.uns["subtype_plan"].columns) == [ + "major", "subtype", "n_cells", "major_type_fraction", "current_k", "expected_k", "budget", + ] + assert "n_subtypes_made" not in result.uns["subtype_plan"].columns + + def test_refine_undersized_true_dispatches_and_populates_subtype_plan(self): + obs = _make_obs({"A": [400, 100]}) # under-clustered vs a forced expected_k=5 + adata = _make_ref_adata(obs, n_genes=20, seed=20) + with pytest.warns(UserWarning, match=r"fewer '.*' labels"): + result = setup_reference( + adata, cell_type_key=["major", "subtype"], max_input_genes=100, + refine_undersized=True, + kws_predefined={"k_heuristic": lambda n: 5}, + kws_refine={"min_frac": 0.0}, + ) + assert list(result.obs.columns) == ["major", "subtype", "subtype_kmeans"] + assert list(result.uns["level_keys"]) == ["major", "subtype", "subtype_kmeans"] + plan = result.uns["subtype_plan"] + assert "n_subtypes_made" in plan.columns + assert set(result.obs["subtype_kmeans"]) == set(f"A_0_{i}" for i in range(4)) | {"A_1_0"} + + def test_n_subtypes_made_matches_actual_realized_output(self): + # Two clearly-separable blocks under one label, forced budget=2 -- + # k-means should cleanly find both, so n_subtypes_made should be 2, + # not just equal to budget by assumption. + block1 = _make_expression(150, 20, seed=30) + block2 = _make_expression(150, 20, seed=31) + block2[:, :5] += 30 + X = np.vstack([block1, block2]) + obs = pd.DataFrame({"major": ["A"] * 300, "subtype": ["A_0"] * 300}) + adata = _make_ref_adata(obs, n_genes=20, seed=21) + adata.X = X + with pytest.warns(UserWarning, match=r"fewer '.*' labels"): + result = setup_reference( + adata, cell_type_key=["major", "subtype"], max_input_genes=100, random_state=0, + refine_undersized=True, + kws_predefined={"k_heuristic": lambda n: 2}, + kws_refine={"min_frac": 0.0}, + ) + plan = result.uns["subtype_plan"].set_index("subtype") + assert plan.loc["A_0", "budget"] == 2 + assert plan.loc["A_0", "n_subtypes_made"] == 2 + + def test_kws_refine_min_frac_reaches_refine_predefined_subtypes(self): + # An extreme min_frac via kws_refine forces every candidate + # sub-cluster to be dropped as noise, collapsing to a single + # fallback subtype despite a budget of 5. + obs = pd.DataFrame({"major": ["A"] * 300, "subtype": ["A_0"] * 300}) + adata = _make_ref_adata(obs, n_genes=20, seed=22) + with pytest.warns(UserWarning) as record: + result = setup_reference( + adata, cell_type_key=["major", "subtype"], max_input_genes=100, random_state=0, + refine_undersized=True, + kws_predefined={"k_heuristic": lambda n: 5}, + kws_refine={"min_frac": 0.99}, + ) + assert len(record) == 2 + assert set(result.obs["subtype_kmeans"]) == {"A_0_0"} + assert result.uns["subtype_plan"].set_index("subtype").loc["A_0", "n_subtypes_made"] == 1 + + def test_refine_undersized_reproducible_with_top_level_random_state(self): + obs = pd.DataFrame({"major": ["A"] * 300, "subtype": ["A_0"] * 300}) + X = _make_expression(300, 20, seed=23) + adata_a = _make_ref_adata(obs, n_genes=20, seed=23) + adata_a.X = X + adata_b = _make_ref_adata(obs, n_genes=20, seed=23) + adata_b.X = X + kwargs = dict( + cell_type_key=["major", "subtype"], max_input_genes=100, random_state=3, + refine_undersized=True, kws_predefined={"k_heuristic": lambda n: 3}, kws_refine={"min_frac": 0.0}, + ) + with pytest.warns(UserWarning, match=r"fewer '.*' labels"): + r_a = setup_reference(adata_a, **kwargs) + with pytest.warns(UserWarning, match=r"fewer '.*' labels"): + r_b = setup_reference(adata_b, **kwargs) + assert list(r_a.obs_names) == list(r_b.obs_names) + + def test_refine_undersized_verbose_prints_subclusters_added(self, capsys): + obs = pd.DataFrame({"major": ["A"] * 300, "subtype": ["A_0"] * 300}) + adata = _make_ref_adata(obs, n_genes=20, seed=24) + with pytest.warns(UserWarning, match=r"fewer '.*' labels"): + setup_reference( + adata, cell_type_key=["major", "subtype"], max_input_genes=100, random_state=0, verbose=True, + refine_undersized=True, + kws_predefined={"k_heuristic": lambda n: 3}, + kws_refine={"min_frac": 0.0}, + ) + out = capsys.readouterr().out + assert re.search(r"Added \d+ subclusters to finest cell type levels", out) + + def test_refine_undersized_verbose_with_categorical_obs_does_not_crash(self, capsys): + # Same categorical-dtype pitfall as + # test_budget_and_n_cells_are_int_with_categorical_labels, but for + # n_subtypes_made: with verbose=True this used to crash with + # "'Categorical' with dtype category does not support operation + # 'sum'" on real AnnData input, since obs columns are categorical + # by default there (unlike this suite's plain-string test helpers). + obs = pd.DataFrame({"major": ["A"] * 300, "subtype": ["A_0"] * 300}) + obs["major"] = obs["major"].astype("category") + obs["subtype"] = obs["subtype"].astype("category") + adata = _make_ref_adata(obs, n_genes=20, seed=26) + with pytest.warns(UserWarning, match=r"fewer '.*' labels"): + result = setup_reference( + adata, cell_type_key=["major", "subtype"], max_input_genes=100, random_state=0, verbose=True, + refine_undersized=True, + kws_predefined={"k_heuristic": lambda n: 2}, + kws_refine={"min_frac": 0.0}, + ) + out = capsys.readouterr().out + assert re.search(r"Added \d+ subclusters to finest cell type levels", out) + assert result.uns["subtype_plan"]["n_subtypes_made"].dtype == np.int64 + + def test_refine_undersized_ignored_on_kmeans_only_path(self): + obs = pd.DataFrame({"cell_type": ["T0"] * 50 + ["T1"] * 50}) + adata = _make_ref_adata(obs, n_genes=20, seed=25) + result = setup_reference( + adata, cell_type_key="cell_type", subcluster_size=3, random_state=0, + max_input_genes=100, refine_undersized=True, + ) + assert "subtype_plan" not in result.uns + assert list(result.obs.columns) == ["cell_type", "subtype"] + + +class TestSymmetricPairsMatrix: + def test_builds_symmetric_matrix_from_one_directional_edges(self): + i_idx = np.array([0, 1]) + j_idx = np.array([1, 2]) + weights = np.array([0.9, 0.7]) + mat = _symmetric_pairs_matrix(i_idx, j_idx, weights, n=4).toarray() + + expected = np.zeros((4, 4)) + expected[0, 1] = expected[1, 0] = 0.9 + expected[1, 2] = expected[2, 1] = 0.7 + assert np.allclose(mat, expected) + + def test_output_is_symmetric(self): + i_idx = np.array([0, 0, 2]) + j_idx = np.array([1, 3, 3]) + weights = np.array([0.5, 0.6, 0.7]) + mat = _symmetric_pairs_matrix(i_idx, j_idx, weights, n=5).toarray() + assert np.allclose(mat, mat.T) + + def test_unconnected_spot_stays_all_zero(self): + i_idx = np.array([0]) + j_idx = np.array([1]) + weights = np.array([1.0]) + mat = _symmetric_pairs_matrix(i_idx, j_idx, weights, n=3).toarray() + assert np.allclose(mat[2, :], 0) + assert np.allclose(mat[:, 2], 0) + + +class TestSetupSpatial: + def test_returns_anndata_with_expected_shape(self): + adata = _make_spatial_adata(15, 8, seed=40) + result = setup_spatial(adata, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0) + assert isinstance(result, AnnData) + assert result.n_obs == 15 + assert result.n_vars == 8 + + def test_obs_preserved_from_input(self): + adata = _make_spatial_adata(10, 5, seed=41) + adata.obs["region"] = ["x"] * 5 + ["y"] * 5 + result = setup_spatial(adata, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0) + assert list(result.obs["region"]) == ["x"] * 5 + ["y"] * 5 + + def test_obsm_spatial_holds_coordinates(self): + adata = _make_spatial_adata(10, 5, seed=42) + coords = adata.obsm["spatial"].copy() + result = setup_spatial(adata, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0) + assert np.allclose(result.obsm["spatial"], coords) + + def test_obsm_spatial_truncates_3d_to_2d(self): + adata = _make_spatial_adata(10, 5, seed=43) + coords_3d = np.hstack([adata.obsm["spatial"], np.zeros((10, 1))]) + adata.obsm["spatial"] = coords_3d + result = setup_spatial(adata, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0) + assert result.obsm["spatial"].shape == (10, 2) + + def test_obsm_spatial_uses_fixed_key_regardless_of_input_spatial_key(self): + adata = _make_spatial_adata(10, 5, seed=44) + coords = adata.obsm.pop("spatial") + adata.obsm["xy"] = coords + result = setup_spatial(adata, spatial_key="xy", th_spatial=0, th_gene_low=0.0, th_gene_high=1.0) + assert "spatial" in result.obsm + assert np.allclose(result.obsm["spatial"], coords) + + def test_gene_frequency_filter_reduces_gene_count(self): + n_spots = 20 + X = np.zeros((n_spots, 3)) + X[:, 0] = 5 # expressed in all spots -> freq 1.0, filtered out + X[:, 1] = 0 # never expressed -> freq 0.0, filtered out + X[:10, 2] = 5 # expressed in half -> freq 0.5, survives + adata = AnnData(X=X, obs=pd.DataFrame(index=[str(i) for i in range(n_spots)])) + adata.var_names = ["always", "never", "half"] + adata.obsm["spatial"] = np.random.default_rng(45).uniform(0, 10, size=(n_spots, 2)) + result = setup_spatial(adata, th_spatial=0) + assert list(result.var_names) == ["half"] + + def test_gene_filter_disabled_when_thresholds_span_full_range(self): + n_spots = 20 + X = np.zeros((n_spots, 3)) + X[:, 0] = 5 + X[:, 1] = 0 + X[:10, 2] = 5 + adata = AnnData(X=X, obs=pd.DataFrame(index=[str(i) for i in range(n_spots)])) + adata.var_names = ["always", "never", "half"] + adata.obsm["spatial"] = np.random.default_rng(46).uniform(0, 10, size=(n_spots, 2)) + result = setup_spatial(adata, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0) + assert list(result.var_names) == ["always", "never", "half"] + + def test_obsp_dot_spatial_pairs_present_when_neighbors_found(self): + adata = _make_clustered_spatial_adata(n_clusters=3, spots_per_cluster=3, seed=47) + result = setup_spatial(adata, radius=2.0, th_gene_low=0.0, th_gene_high=1.0) + assert "dot_spatial_pairs" in result.obsp + mat = result.obsp["dot_spatial_pairs"] + assert mat.shape == (9, 9) + assert mat.nnz > 0 + + def test_obsp_dot_spatial_pairs_symmetric(self): + adata = _make_clustered_spatial_adata(n_clusters=3, spots_per_cluster=3, seed=48) + result = setup_spatial(adata, radius=2.0, th_gene_low=0.0, th_gene_high=1.0) + mat = result.obsp["dot_spatial_pairs"] + assert np.allclose(mat.toarray(), mat.toarray().T) + + def test_obsp_weights_are_cosine_similarities(self): + adata = _make_clustered_spatial_adata(n_clusters=3, spots_per_cluster=3, seed=49) + result = setup_spatial(adata, radius=2.0, th_gene_low=0.0, th_gene_high=1.0) + dense = result.obsp["dot_spatial_pairs"].toarray() + nonzero = dense[dense != 0] + assert nonzero.size > 0 + # identical intra-cluster profiles -> cosine similarity exactly 1.0 + assert np.allclose(nonzero, 1.0) + + def test_obsp_dot_spatial_pairs_omitted_when_th_spatial_disabled(self): + adata = _make_clustered_spatial_adata(n_clusters=3, spots_per_cluster=3, seed=50) + result = setup_spatial(adata, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0) + assert "dot_spatial_pairs" not in result.obsp + + def test_obsp_dot_spatial_pairs_omitted_when_no_pairs_found(self): + adata = _make_clustered_spatial_adata(n_clusters=3, spots_per_cluster=3, seed=51) + result = setup_spatial(adata, th_spatial=1.5, th_gene_low=0.0, th_gene_high=1.0) + assert "dot_spatial_pairs" not in result.obsp + + def test_radius_auto_estimates_from_data(self): + adata = _make_clustered_spatial_adata(n_clusters=2, spots_per_cluster=10, seed=52) + result = setup_spatial(adata, radius="auto", th_gene_low=0.0, th_gene_high=1.0) + assert "dot_spatial_pairs" in result.obsp + + def test_radius_numeric_larger_finds_at_least_as_many_pairs(self): + n = 10 + coords = np.column_stack([np.arange(n, dtype=float), np.zeros(n)]) + base = _make_expression(1, 5, seed=53) + X = np.repeat(base, n, axis=0) + adata = AnnData(X=X, obs=pd.DataFrame(index=[str(i) for i in range(n)])) + adata.var_names = [f"gene{i}" for i in range(5)] + adata.obsm["spatial"] = coords + + r_small = setup_spatial(adata, radius=1.5, th_gene_low=0.0, th_gene_high=1.0) + r_large = setup_spatial(adata, radius=3.5, th_gene_low=0.0, th_gene_high=1.0) + n_small = r_small.obsp["dot_spatial_pairs"].nnz if "dot_spatial_pairs" in r_small.obsp else 0 + n_large = r_large.obsp["dot_spatial_pairs"].nnz if "dot_spatial_pairs" in r_large.obsp else 0 + assert n_small > 0 + assert n_large >= n_small + + def test_missing_spatial_key_raises(self): + adata = AnnData(X=_make_expression(5, 4, seed=54), obs=pd.DataFrame(index=[str(i) for i in range(5)])) + with pytest.raises(ValueError, match="obsm"): + setup_spatial(adata) + + def test_warn_mt_genes_warns_when_present(self): + adata = _make_spatial_adata(10, 5, seed=55) + genes = [f"gene{i}" for i in range(5)] + genes[0] = "MT-ND1" + adata.var_names = genes + with pytest.warns(UserWarning, match="MT-"): + setup_spatial(adata, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0) + + def test_warn_mt_false_silences_warning(self, recwarn): + adata = _make_spatial_adata(10, 5, seed=56) + genes = [f"gene{i}" for i in range(5)] + genes[0] = "MT-ND1" + adata.var_names = genes + setup_spatial(adata, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0, warn_mt=False) + assert len(recwarn) == 0 + + def test_copy_true_does_not_mutate_input(self): + adata = _make_spatial_adata(10, 5, seed=57) + X_before = adata.X.copy() + setup_spatial(adata, copy=True, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0) + assert np.allclose(adata.X, X_before) + + def test_sparse_X_works_end_to_end(self): + adata_dense = _make_spatial_adata(10, 5, seed=58) + adata_sparse = adata_dense.copy() + adata_sparse.X = csr_matrix(adata_dense.X) + r_dense = setup_spatial(adata_dense, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0) + r_sparse = setup_spatial(adata_sparse, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0) + X_sparse_dense = r_sparse.X.toarray() if issparse(r_sparse.X) else r_sparse.X + assert np.allclose(r_dense.X, X_sparse_dense) + + def test_verbose_prints_progress(self, capsys): + adata = _make_spatial_adata(10, 5, seed=59) + setup_spatial(adata, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0, verbose=True) + out = capsys.readouterr().out + assert out != "" + + def test_verbose_false_prints_nothing(self, capsys): + adata = _make_spatial_adata(10, 5, seed=60) + setup_spatial(adata, th_spatial=0, th_gene_low=0.0, th_gene_high=1.0, verbose=False) + out = capsys.readouterr().out + assert out == "" + + def test_h5ad_round_trip_preserves_everything(self, tmp_path): + adata = _make_clustered_spatial_adata(n_clusters=3, spots_per_cluster=3, seed=61) + result = setup_spatial(adata, radius=2.0, th_gene_low=0.0, th_gene_high=1.0) + path = tmp_path / "spatial.h5ad" + result.write_h5ad(path) + loaded = sc.read_h5ad(path) + + assert np.allclose(np.asarray(loaded.X), np.asarray(result.X)) + assert list(loaded.var_names) == list(result.var_names) + assert np.allclose(loaded.obsm["spatial"], result.obsm["spatial"]) + assert np.allclose( + loaded.obsp["dot_spatial_pairs"].toarray(), result.obsp["dot_spatial_pairs"].toarray() + ) + + def test_th_nonspatial_augments_pairs_beyond_spatial_radius(self): + # 6 clusters of 2 spots each, all sharing the same expression + # profile but spaced far enough apart that only the 6 within + # -cluster pairs are found spatially -- th_nonspatial should add + # cross-cluster pairs (gene-expression-similar despite being out + # of radius) up to the max_pairs=N cap. + rng = np.random.default_rng(0) + n_genes = 8 + n_clusters = 6 + base = rng.integers(1, 10, size=(1, n_genes)).astype(np.float64) + + X_parts, coords_parts = [], [] + for c in range(n_clusters): + X_parts.append(base + rng.normal(0, 0.01, size=(2, n_genes))) + coords_parts.append(rng.uniform(-0.3, 0.3, size=(2, 2)) + np.array([c * 1000.0, 0.0])) + X = np.clip(np.vstack(X_parts), 0.1, None) + coords = np.vstack(coords_parts) + N = X.shape[0] + + adata = AnnData(X=X, obs=pd.DataFrame(index=[str(i) for i in range(N)])) + adata.var_names = [f"gene{i}" for i in range(n_genes)] + adata.obsm["spatial"] = coords + + result_off = setup_spatial( + adata.copy(), radius=2.0, th_gene_low=0.0, th_gene_high=1.0, th_nonspatial=0.0, + ) + result_on = setup_spatial( + adata.copy(), radius=2.0, th_gene_low=0.0, th_gene_high=1.0, th_nonspatial=0.9, + ) + n_pairs_off = result_off.obsp["dot_spatial_pairs"].nnz // 2 + n_pairs_on = result_on.obsp["dot_spatial_pairs"].nnz // 2 + assert n_pairs_off == n_clusters # one pair per cluster + assert n_pairs_on > n_pairs_off + assert n_pairs_on <= N # capped at max_pairs + + def test_sparse_X_pairwise_cosine_matches_dense(self): + # test_sparse_X_works_end_to_end above uses th_spatial=0, which + # skips pair-finding entirely -- this specifically exercises the + # sparse branch of the pairwise cosine-similarity computation + # inside the radius-neighbor pair search. + adata_dense = _make_clustered_spatial_adata(n_clusters=3, spots_per_cluster=3, seed=62) + adata_sparse = adata_dense.copy() + adata_sparse.X = csr_matrix(adata_dense.X) + + r_dense = setup_spatial(adata_dense, radius=2.0, th_gene_low=0.0, th_gene_high=1.0) + r_sparse = setup_spatial(adata_sparse, radius=2.0, th_gene_low=0.0, th_gene_high=1.0) + + assert np.allclose( + r_dense.obsp["dot_spatial_pairs"].toarray(), r_sparse.obsp["dot_spatial_pairs"].toarray(), + ) diff --git a/tests/test_visualization.py b/tests/test_visualization.py new file mode 100644 index 0000000..00b9123 --- /dev/null +++ b/tests/test_visualization.py @@ -0,0 +1,146 @@ +import numpy as np +import pandas as pd + +from dotpy.visualization import ( + plot_cell_type_proportions, + plot_optimization_history, + plot_spatial_weights, +) + + +def _make_weights(n_spots=6, n_types=3, seed=0): + rng = np.random.default_rng(seed) + return rng.uniform(0.1, 1.0, size=(n_spots, n_types)) + + +def _make_coords(n_spots=6, seed=0): + rng = np.random.default_rng(seed) + return rng.uniform(0, 10, size=(n_spots, 2)) + + +class TestPlotSpatialWeights: + def test_runs_with_ndarray_weights(self): + fig = plot_spatial_weights(_make_coords(), _make_weights(), cell_types=["A", "B", "C"]) + assert fig is not None + + def test_runs_with_dataframe_weights(self): + weights = pd.DataFrame(_make_weights(), columns=["A", "B", "C"]) + fig = plot_spatial_weights(_make_coords(), weights, cell_types=["A", "B", "C"]) + assert fig is not None + + def test_dataframe_and_ndarray_produce_same_plotted_values(self): + coords = _make_coords() + weights_arr = _make_weights() + weights_df = pd.DataFrame(weights_arr, columns=["A", "B", "C"]) + + fig_arr = plot_spatial_weights(coords, weights_arr, cell_types=["A", "B", "C"]) + fig_df = plot_spatial_weights(coords, weights_df, cell_types=["A", "B", "C"]) + + for i in range(weights_arr.shape[1]): + scatter_arr = fig_arr.axes[i].collections[0] + scatter_df = fig_df.axes[i].collections[0] + np.testing.assert_allclose(scatter_arr.get_array(), scatter_df.get_array()) + np.testing.assert_allclose(scatter_arr.get_offsets(), scatter_df.get_offsets()) + + def test_single_cell_type_with_ncols_one_does_not_raise(self): + # nrows == ncols == 1 is the only case where plt.subplots() returns a + # bare Axes instead of an array, which plot_spatial_weights has to + # special-case. + weights = _make_weights(n_types=1) + fig = plot_spatial_weights(_make_coords(), weights, cell_types=["OnlyType"], ncols=1) + assert len(fig.axes[0].collections) == 1 + + def test_zero_row_sum_does_not_produce_nan(self): + coords = _make_coords(n_spots=4) + weights = _make_weights(n_spots=4, n_types=2) + weights[0, :] = 0.0 + fig = plot_spatial_weights(coords, weights, cell_types=["A", "B"], normalize=True) + for ax in fig.axes[:2]: + assert not np.any(np.isnan(ax.collections[0].get_array())) + + def test_default_cell_type_labels(self): + weights = _make_weights(n_types=2) + fig = plot_spatial_weights(_make_coords(), weights) + titles = [ax.get_title() for ax in fig.axes[:2]] + assert titles == ["CT1", "CT2"] + + def test_save_path_writes_file(self, tmp_path): + save_path = tmp_path / "plot.png" + plot_spatial_weights( + _make_coords(), _make_weights(), cell_types=["A", "B", "C"], save_path=str(save_path) + ) + assert save_path.exists() + + def test_title_is_set_as_figure_suptitle(self): + fig = plot_spatial_weights( + _make_coords(), _make_weights(), cell_types=["A", "B", "C"], title="My Title", + ) + assert fig._suptitle.get_text() == "My Title" + + +class TestPlotCellTypeProportions: + def test_runs_with_dataframe_weights(self): + weights = pd.DataFrame(_make_weights(), columns=["A", "B", "C"]) + fig = plot_cell_type_proportions(weights, cell_types=["A", "B", "C"]) + assert fig is not None + + def test_default_cell_type_labels(self): + fig = plot_cell_type_proportions(_make_weights(n_types=2)) + ax = fig.axes[0] + labels = {t.get_text() for t in ax.get_yticklabels()} + assert labels == {"CT1", "CT2"} + + def test_sorted_labels_match_sorted_bars(self): + # Column "B" >> "C" >> "A" in total weight, so after sorting, labels + # and bar widths should both read B, C, A in lockstep. + weights = np.array([ + [0.1, 5.0, 0.6], + [0.1, 5.0, 0.6], + ]) + fig = plot_cell_type_proportions(weights, cell_types=["A", "B", "C"]) + ax = fig.axes[0] + labels = [t.get_text() for t in ax.get_yticklabels()] + widths = [bar.get_width() for bar in ax.patches] + assert labels == ["B", "C", "A"] + assert widths == sorted(widths, reverse=True) + + def test_save_path_writes_file(self, tmp_path): + save_path = tmp_path / "proportions.png" + plot_cell_type_proportions( + _make_weights(), cell_types=["A", "B", "C"], save_path=str(save_path) + ) + assert save_path.exists() + + +def _make_history(n=5, all_none_lower_bound=False): + return { + "iteration": list(range(1, n + 1)), + "objective": list(np.linspace(10, 5, n)), + "upper_bound": list(np.linspace(10, 5, n)), + "lower_bound": [None] * n if all_none_lower_bound else list(np.linspace(3, 4, n)), + "gap": list(np.linspace(0.5, 0.01, n)), + "time": list(np.linspace(0.1, 0.1, n)), + } + + +class TestPlotOptimizationHistory: + def test_runs_and_returns_figure(self): + fig = plot_optimization_history(_make_history()) + assert fig is not None + assert len(fig.axes) == 3 + + def test_save_path_writes_file(self, tmp_path): + save_path = tmp_path / "history.png" + plot_optimization_history(_make_history(), save_path=str(save_path)) + assert save_path.exists() + + def test_all_none_lower_bound_omits_lower_bound_line(self): + # history['lower_bound'] is all None early in a fit (before the + # first duality-gap improvement) -- the plot should skip that + # line rather than plotting a line of NaNs. + fig = plot_optimization_history(_make_history(all_none_lower_bound=True)) + assert len(fig.axes[0].lines) == 2 + + def test_normal_lower_bound_includes_lower_bound_line(self): + fig = plot_optimization_history(_make_history(all_none_lower_bound=False)) + assert len(fig.axes[0].lines) == 3