Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
64 commits
Select commit Hold shift + click to select a range
69e419e
Add find_stack_level() helper for warning attribution
robinfallegger Aug 7, 2026
c14448a
moved preprocessing steps out of the setup_reference function, and re…
robinfallegger Aug 6, 2026
4b70adf
setup_spatial: changed MT/HLA/RPL removal to warning (up to user to r…
robinfallegger Aug 6, 2026
f90516a
added non-negative check skip option
robinfallegger Aug 6, 2026
cf7fe21
use automatic stack level finder for warnings
robinfallegger Aug 7, 2026
7541ca6
make reference input validation function public
robinfallegger Aug 7, 2026
ec642b0
Add .gitignore for Python build/cache artifacts
robinfallegger Aug 7, 2026
97d086b
Fix _is_in_package() to not require Python 3.9's Path.is_relative_to()
robinfallegger Aug 7, 2026
0548f12
Add pytest test harness (pytest.ini, requirements-test.txt)
robinfallegger Aug 7, 2026
17af0e7
Add toy_reference() synthetic dataset for future testing
robinfallegger Aug 7, 2026
b70e4a4
Add tests for toy_reference()
robinfallegger Aug 7, 2026
26fca01
function to check for unique nested cell type hierarchy
robinfallegger Aug 7, 2026
49016e3
tests for check_subtype_consistency
robinfallegger Aug 7, 2026
5aa6351
Add kmeans_define_subtypes() and expose min_frac, k_heuristic, max_ce…
robinfallegger Aug 7, 2026
d4704f8
Add tests for kmeans_define_subtypes()
robinfallegger Aug 7, 2026
4e8fbb3
Add predefined_subtypes() for adapting and checking user-supplied mu…
robinfallegger Aug 7, 2026
53021ac
Add tests for predefined_subtypes()
robinfallegger Aug 7, 2026
a4fb2e2
Change kmeans_define_subtypes() to return a DataFrame
robinfallegger Aug 7, 2026
b83885c
Add summarize_subtypes() to aggregate cells into per-subtype centroid…
robinfallegger Aug 8, 2026
61b137f
Add tests for summarize_subtypes()
robinfallegger Aug 8, 2026
c738bee
Add equivalence tests comparing the new pipeline to _aggregate_reference
robinfallegger Aug 8, 2026
100a6d3
Add select_de_genes() with pluggable scoring methods
robinfallegger Aug 10, 2026
53f3b72
Add tests for select_de_genes()
robinfallegger Aug 10, 2026
9afa597
Add equivalence tests comparing select_de_genes() to _get_de_genes_r_…
robinfallegger Aug 10, 2026
2dc533d
Wire predefined/k-means subtype discovery into setup_reference(); ret…
robinfallegger Aug 10, 2026
a3ed1e6
Add tests for (new) setup_reference()
robinfallegger Aug 10, 2026
6654876
Add equivalence tests comparing setup_reference() to the legacy pipeline
robinfallegger Aug 10, 2026
0c2a26f
Add uns['level_keys'] to setup_reference()'s AnnData output
robinfallegger Aug 10, 2026
b1dcf46
Remove _aggregate_reference and _get_de_genes_r_style
robinfallegger Aug 10, 2026
24e1920
Rework setup_spatial() to return AnnData instead of a dict
robinfallegger Aug 10, 2026
0d69c6c
Add tests for setup_spatial()
robinfallegger Aug 10, 2026
08740dc
Use AnnData for ref/spatial in DOT.__init__(); future warning for dict
robinfallegger Aug 10, 2026
e99b856
moved common testing helper functions to new file
robinfallegger Aug 10, 2026
f409139
Add tests for DOT.__init__() and its AnnData/dict conversion helpers
robinfallegger Aug 10, 2026
e58de7c
Add level= support to get_weights()/get_cell_types(); return DataFrame
robinfallegger Aug 11, 2026
8149fdc
Add tests for get_weights()/get_cell_types() level= support
robinfallegger Aug 11, 2026
42bc23d
Move optimisation config to DOT construction time; add direct lambda …
robinfallegger Aug 11, 2026
e6b3538
Add tests for the optimisation weight/lambda config system
robinfallegger Aug 11, 2026
62da2b7
Add print_config() to DOT to show optimisation weights
robinfallegger Aug 11, 2026
a68e229
Add tests for print_config()
robinfallegger Aug 11, 2026
eed874a
Accept DataFrame or ndarray weights in plotting functions
robinfallegger Aug 11, 2026
ebdbbcd
Add tests for plot_spatial_weights and plot_cell_type_proportions
robinfallegger Aug 11, 2026
09971df
Update docs and scripts for new DOT()/fit() config API
robinfallegger Aug 11, 2026
0893616
Add cluster_weight/l_c to API. Not implemented in optimsation yet
robinfallegger Aug 11, 2026
c13af6e
Adapt/add tests for cluster_weight/l_c config wiring
robinfallegger Aug 11, 2026
1c24080
Implement l_c (cluster-wise cosine) in the optimisation loop
robinfallegger Aug 11, 2026
0108f22
Add tests for cluster_weight/l_c optimisation behaviour
robinfallegger Aug 11, 2026
e336c28
Batch the cluster-wise cosine term for memory-bounded execution
robinfallegger Aug 11, 2026
5fa44e6
Add tests for batch cluster-cosine optimisation
robinfallegger Aug 11, 2026
7cd0d29
Fix _safe_log2 to match R's safelog2 exactly
robinfallegger Aug 11, 2026
b497fca
Add tests for _sqrt_env/_sqrt_env_grad envelope
robinfallegger Aug 11, 2026
3140142
Add tests for validate_reference_input and _de_scores_logfc_rank
robinfallegger Aug 11, 2026
8320053
Close test-coverage gaps found via pytest-cov (87% -> 96%)
robinfallegger Aug 11, 2026
5568386
Add plan_subtype_refinement + apportionment for under-clustered subty…
robinfallegger Aug 12, 2026
bf3219b
Add tests for _expected_k, _huntington_hill_apportion, plan_subtype_r…
robinfallegger Aug 12, 2026
66cf13e
use subtype refinement in predefined_subtypes (for warnings and count…
robinfallegger Aug 12, 2026
d2424db
Compute _kmeans_subcluster's min_cells threshold at the call site ins…
robinfallegger Aug 12, 2026
17ebacc
extract sampled group centroid calculation into _capped_group_centroi…
robinfallegger Aug 13, 2026
65cfc1f
Add refine_predefined_subtypes: budget-driven sub-clustering of under…
robinfallegger Aug 13, 2026
80c0d09
add tests for refine_predefined_subtypes
robinfallegger Aug 13, 2026
bebe9a0
wire refine_undersized/kws_refine into setup_reference; consolidate p…
robinfallegger Aug 14, 2026
abf9f43
add tests for setup_reference refinement wiring and aggregate warnings
robinfallegger Aug 14, 2026
281827a
fixed type error in plan_subtype_refinement
robinfallegger Aug 25, 2026
511f38b
added changelog
robinfallegger Aug 25, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 34 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
@@ -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
34 changes: 34 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -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.
66 changes: 46 additions & 20 deletions QUICKSTART.md
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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)
Expand All @@ -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!")
Expand All @@ -143,24 +149,37 @@ 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
)
```

**When to adjust:**
- 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

Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
55 changes: 32 additions & 23 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -96,7 +103,6 @@ plot_spatial_weights(

```python
dot.fit(
mode='highres',
iterations=100,
resume_from='./checkpoints/checkpoint_iter_50.pkl',
verbose=True
Expand Down Expand Up @@ -245,26 +251,28 @@ 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)

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
Expand All @@ -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
)

Expand All @@ -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
)

Expand All @@ -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
Expand All @@ -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')
Expand Down Expand Up @@ -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
)
Expand All @@ -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
)
Expand Down Expand Up @@ -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')
Expand Down
Loading