ReMEDi / scripts /align_genes.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw
History Blame Contribute Delete
1.72 kB
"""Align raw study cohorts to shared gene identifiers before fitting PCA."""
import argparse
from pathlib import Path
import anndata as ad
p = argparse.ArgumentParser()
p.add_argument('--inputs', nargs='+', required=True)
p.add_argument('--output', required=True)
p.add_argument('--gene-columns', nargs='+', help='One var column per input. Use index for var_names.')
p.add_argument('--drop-ambiguous', action='store_true', help='Drop all duplicated gene identifiers')
a = p.parse_args()
cohorts = [ad.read_h5ad(f) for f in a.inputs]
if a.gene_columns and len(a.gene_columns) != len(cohorts): raise ValueError('Supply one gene-column choice per input')
for i, cohort in enumerate(cohorts):
field = a.gene_columns[i] if a.gene_columns else 'index'
if field != 'index':
if cohort.var[field].isna().any():
cohort = cohort[:, cohort.var[field].notna()].copy()
cohort.var_names = cohort.var[field].astype(str)
if not cohort.var_names.is_unique:
if not a.drop_ambiguous: raise ValueError('Duplicate gene identifiers. Use --drop-ambiguous to discard these genes.')
cohort = cohort[:, ~cohort.var_names.duplicated(keep=False)].copy()
cohorts[i] = cohort
shared = sorted(set.intersection(*(set(c.var_names) for c in cohorts)))
if len(shared) < 100:
raise ValueError('Fewer than 100 shared genes. Check identifier conventions.')
out = Path(a.output); out.mkdir(parents=True, exist_ok=True)
for i, (path, cohort) in enumerate(zip(a.inputs, cohorts)):
cohort[:, shared].copy().write_h5ad(out / f'{i}_{Path(path).stem}.h5ad', compression='gzip')
(out / 'genes.txt').write_text('\n'.join(shared)+'\n')
print(f'Aligned {len(cohorts)} cohorts on {len(shared)} genes')