diff --git a/CHANGELOG.md b/CHANGELOG.md index fb3d361..47d2274 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,8 @@ * Added `spaTrack` method (PR #4). * Added `Spearman's correlation` metric (PR #5). +* Added `stlearn`method (PR #9). + ## MAJOR CHANGES * Updated `api` files and set the data processor (PR #1). diff --git a/src/methods/stlearn/config.vsh.yaml b/src/methods/stlearn/config.vsh.yaml new file mode 100644 index 0000000..b76f797 --- /dev/null +++ b/src/methods/stlearn/config.vsh.yaml @@ -0,0 +1,82 @@ +__merge__: ../../api/comp_method.yaml + +name: stlearn +label: stLearn +summary: "stLearn reconstructs spatial trajectories by combining diffusion pseudotime on the gene expression with the spatial arrangement of the annotated clusters." +description: | + stLearn infers a "pseudo-time-space" (PSTS) trajectory. Cells are first embedded with PCA and + connected in a k-NN graph, after which every annotated cell type is split into spatially + contiguous sub-clusters with DBSCAN. A PAGA graph over these clusters is combined with + diffusion pseudotime (DPT) to order the cells, and the resulting graph is oriented by + pseudotime to enumerate the trajectories that start at the root cluster. + + The root cluster is determined automatically from the data: the cell type whose cells express + the largest number of genes on average is taken to be the least differentiated one. This is the + same CytoTRACE-like proxy that stLearn uses internally to select the root cell within the root + cluster. Cells that do not lie on any trajectory starting from the root are reported as NaN. +references: + doi: + - 10.1038/s41467-023-43120-6 +links: + documentation: https://stlearn.readthedocs.io/en/latest/ + repository: https://github.com/BiomedicalMachineLearning/stLearn + + + +# Metadata for your component +info: + preferred_normalization: log_cp10k + +arguments: + - name: "--n_comps" + type: "integer" + default: 50 + description: Number of principal components to compute. + - name: "--n_neighbors" + type: "integer" + default: 200 + description: Number of neighbors used to build the k-NN graph. + - name: "--resolution" + type: "double" + default: 0.8 + description: Resolution of the Leiden clustering. + - name: "--eps" + type: "double" + default: 1500 + description: | + Maximum distance between two spots for them to be considered spatial neighbours by the + DBSCAN sub-clustering, in the units of the spatial coordinates. + - name: "--seed" + type: "integer" + default: 0 + description: Random seed. + +resources: + - type: python_script + path: script.py + +engines: + # custom image because stlearn pins numpy<2 and requires python >=3.10,<3.13 + - type: docker + image: python:3.11-slim + setup: + - type: apt + packages: + - procps # required by Nextflow + - git # pip needs it to install openproblems core from git+https + - build-essential # compiler for any source builds + - type: python + upgrade: true + github: + - "openproblems-bio/core#subdirectory=packages/python/openproblems" + packages: + - pyyaml + - requests + - jsonschema + - stlearn==1.2.2 + +runners: + - type: executable + - type: nextflow + directives: + label: [midtime,midmem,midcpu] diff --git a/src/methods/stlearn/script.py b/src/methods/stlearn/script.py new file mode 100644 index 0000000..dde7002 --- /dev/null +++ b/src/methods/stlearn/script.py @@ -0,0 +1,157 @@ +import os +import random +import warnings + +# stlearn imports tensorflow, which logs its device setup to stderr on import +os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3' + +import anndata as ad +import numpy as np +import pandas as pd +import scanpy as sc +import scipy.sparse as sp +import stlearn as st + +## VIASH START +par = { + 'input': 'resources_test/task_spatial_trajectory_inference/dlpfc_151673/dataset.h5ad', + 'output': 'output.h5ad', + 'n_comps': 50, + 'n_neighbors': 200, + 'resolution': 0.8, + 'eps': 1500.0, + 'seed': 0, +} +meta = { + 'name': 'stlearn' +} +## VIASH END + +# warnings raised by stlearn internals, not actionable from here +warnings.simplefilter('ignore', FutureWarning) +warnings.simplefilter('ignore', ad.ImplicitModificationWarning) + +seed = par['seed'] +np.random.seed(seed) +random.seed(seed) + + +def build_cell_type_int(adata): + """Map cell_type strings to integer labels, as stLearn expects numeric cluster labels.""" + adata.obs['cell_type'] = adata.obs['cell_type'].astype(str).astype('category') + unique_cell_types = adata.obs['cell_type'].cat.categories + + ct_to_num = {str(ct): str(i) for i, ct in enumerate(unique_cell_types)} + num_to_ct = {str(i): str(ct) for i, ct in enumerate(unique_cell_types)} + + adata.obs['cell_type_int'] = ( + adata.obs['cell_type'].astype(str).map(ct_to_num).astype('category') + ) + return num_to_ct + + +def select_root(adata): + """Return the least differentiated cluster, i.e. the one whose cells express the largest + number of genes on average. This is the same proxy stLearn uses to pick the root cell.""" + n_expressed = np.asarray((adata.layers['counts'] > 0).sum(axis=1)).reshape(-1) + scores = ( + pd.DataFrame({'n_expressed': n_expressed, 'cluster': adata.obs['cell_type_int'].values}) + .groupby('cluster', observed=True)['n_expressed'] + .mean() + ) + return str(scores.idxmax()) + + +def filter_branches(available_paths, root): + """Keep the paths starting from the root and drop branches contained in a longer one.""" + valid = [path for path in available_paths.values() if path[0] == root] + valid.sort(key=len, reverse=True) + + unique = [] + for branch in valid: + if not any(set(branch).issubset(set(kept)) for kept in unique): + unique.append(branch) + return unique + + +print('Reading input files', flush=True) +adata = ad.read_h5ad(par['input']) + +adata.X = adata.layers['normalized'] +if sp.issparse(adata.X): + adata.X = adata.X.toarray() + +# stLearn reads the spatial coordinates from these slots +adata.obsm['spatial'] = adata.obsm['X_spatial'] +adata.obs['imagerow'] = adata.obsm['X_spatial'][:, 1] +adata.obs['imagecol'] = adata.obsm['X_spatial'][:, 0] + +print('Embed and cluster the cells', flush=True) +st.em.run_pca(adata, n_comps=par['n_comps']) +sc.pp.neighbors(adata, n_neighbors=par['n_neighbors'], use_rep='X_pca') +st.tl.clustering.leiden(adata, resolution=par['resolution'], random_state=seed) + +num_to_ct = build_cell_type_int(adata) + +print('Determine root cell', flush=True) +root = select_root(adata) +print(f'Root cluster: {num_to_ct[root]}', flush=True) + +# use_raw is False because the dataset does not carry a raw layer +adata.uns['iroot'] = st.spatial.trajectory.set_root( + adata, + use_label='cell_type_int', + cluster=int(root), + use_raw=False, +) + +print('Calculate the pseudotime and the trajectory branches', flush=True) +st.spatial.trajectory.pseudotime( + adata, + eps=par['eps'], + use_rep='X_pca', + use_label='cell_type_int', +) +branches = filter_branches(adata.uns.get('available_paths', {}), int(root)) + +if not branches: + raise RuntimeError( + f'No valid branches found from root cluster {root}. Try adjusting --resolution or --eps.' + ) + +print('Generate predictions', flush=True) +# cells not covered by any branch keep a NaN pseudotime +adata.obs['pseudotime_inferred'] = np.nan + +for branch in branches: + branch_labels = [str(node) for node in branch] + try: + st.spatial.trajectory.pseudotimespace_global( + adata, + use_label='cell_type_int', + list_clusters=branch_labels, + ) + except Exception as exc: + print(f'Skipping branch {branch}: {exc}', flush=True) + continue + + # the first branch a cell belongs to takes precedence + unfilled = ( + adata.obs['cell_type_int'].isin(branch_labels) + & adata.obs['pseudotime_inferred'].isna() + ) + adata.obs.loc[unfilled, 'pseudotime_inferred'] = adata.obs.loc[unfilled, 'dpt_pseudotime'] + +n_assigned = adata.obs['pseudotime_inferred'].notna().sum() +print(f'Pseudotime assigned to {n_assigned}/{adata.n_obs} cells', flush=True) + +print('Write output AnnData to file', flush=True) +output = ad.AnnData( + obs=adata.obs[['pseudotime_inferred']], + uns={ + 'dataset_id': adata.uns['dataset_id'], + 'normalization_id': adata.uns['normalization_id'], + 'method_id': meta['name'], + }, +) +output.write_h5ad(par['output'], compression='gzip')