Download scripts/prepare_tahoe.py from ChatterjeeLab/ReMEDi: direct link, hf CLI and curl.
- Browser
- Download file 4.29 kB
-
https://huggingface.co/ChatterjeeLab/ReMEDi/resolve/main/scripts/prepare_tahoe.py
- Command line
-
hf download hf://ChatterjeeLab/ReMEDi/scripts/prepare_tahoe.py
-
curl -L -o prepare_tahoe.py https://huggingface.co/ChatterjeeLab/ReMEDi/resolve/main/scripts/prepare_tahoe.py
4.29 kB
| """Convert official Tahoe expression shards to a bounded raw-count AnnData cohort.""" | |
| import argparse | |
| import ast | |
| from collections import Counter | |
| from pathlib import Path | |
| import json | |
| import anndata as ad | |
| import numpy as np | |
| import pandas as pd | |
| import pyarrow.parquet as pq | |
| from scipy import sparse | |
| from remedi.data import UNIT_TO_UM | |
| def concentration(value, drug): | |
| terms = ast.literal_eval(value) | |
| if len(terms) != 1: | |
| raise ValueError('Only single-molecule perturbations are supported') | |
| name, dose, unit = terms[0] | |
| if name.strip() != drug.strip(): | |
| raise ValueError('Sample compound differs from expression compound') | |
| return float(dose) * UNIT_TO_UM[unit] | |
| def convert(shards, metadata, output, cell_line=None, cap=128, seed=0): | |
| meta = Path(metadata) | |
| samples = pd.DataFrame(pq.read_table(meta/'sample_metadata.parquet').to_pylist()).set_index('sample') | |
| genes = pd.DataFrame(pq.read_table(meta/'gene_metadata.parquet').to_pylist()) | |
| gene_ids = genes.ensembl_id.astype(str).tolist() | |
| if len(set(gene_ids)) != len(gene_ids): raise ValueError('Duplicate gene identifiers') | |
| token = dict(zip(genes.token_id.astype(int), range(len(genes)))) | |
| rng = np.random.default_rng(seed) | |
| pools, seen = {}, Counter() | |
| for path in shards: | |
| for batch in pq.ParquetFile(path).iter_batches(batch_size=1024): | |
| for row in batch.to_pylist(): | |
| if cell_line and row['cell_line_id'] != cell_line: continue | |
| sample = samples.loc[row['sample']] | |
| if isinstance(sample, pd.DataFrame): raise ValueError('Duplicate sample identifiers') | |
| if sample.plate != row['plate']: raise ValueError('Plate metadata mismatch') | |
| dose = 0. if row['drug']=='DMSO_TF' else concentration(sample.drugname_drugconc, row['drug']) | |
| key = (row['drug'], dose, row['cell_line_id'], row['plate']) | |
| seen[key] += 1 | |
| pool = pools.setdefault(key, []) | |
| j = len(pool) if len(pool) < cap else int(rng.integers(seen[key])) | |
| if j >= cap: continue | |
| # The first gene/expression pair is the released CLS token. | |
| ids, values = row['genes'][1:], row['expressions'][1:] | |
| if len(ids) != len(values): raise ValueError('Misaligned genes and counts') | |
| cols = [token[int(i)] for i in ids] | |
| record = ({k: row[k] for k in ['drug','sample','cell_line_id','plate','canonical_smiles','BARCODE_SUB_LIB_ID']}, cols, values, dose) | |
| if j == len(pool): pool.append(record) | |
| else: pool[j] = record | |
| print(f'Read {path}. Retained {sum(map(len,pools.values()))} cells', flush=True) | |
| observations, columns, values, pointers = [], [], [], [0] | |
| for key in sorted(pools): | |
| for row, cols, vals, dose in pools[key]: | |
| row['dose_um'] = dose | |
| observations.append(row); columns.extend(cols); values.extend(vals); pointers.append(len(values)) | |
| if not observations: raise ValueError('No matching cells') | |
| obs = pd.DataFrame(observations) | |
| obs.index = obs.BARCODE_SUB_LIB_ID.astype(str) | |
| if not obs.index.is_unique: raise ValueError('Duplicate source cells across shards') | |
| x = sparse.csr_matrix((np.asarray(values, dtype=np.float32), columns, pointers), shape=(len(obs),len(genes))) | |
| out = Path(output); out.mkdir(parents=True,exist_ok=True) | |
| ad.AnnData(x, obs=obs, var=genes.set_index('ensembl_id')).write_h5ad(out/'tahoe.h5ad',compression='gzip') | |
| obs.loc[obs.drug!='DMSO_TF',['drug','canonical_smiles']].drop_duplicates().rename(columns={'canonical_smiles':'smiles'}).to_csv(out/'structures.csv',index=False) | |
| (out/'source.json').write_text(json.dumps({'dataset':'tahoebio/Tahoe-100M','shards':[str(p) for p in shards], 'cells':len(obs),'conditions':len(pools),'reservoir_cap':cap,'seed':seed},indent=2)+'\n') | |
| if __name__ == '__main__': | |
| p=argparse.ArgumentParser() | |
| p.add_argument('--shards',nargs='+',required=True) | |
| p.add_argument('--metadata',required=True) | |
| p.add_argument('--output',required=True) | |
| p.add_argument('--cell-line') | |
| p.add_argument('--cap',type=int,default=128) | |
| p.add_argument('--seed',type=int,default=0) | |
| convert(**vars(p.parse_args())) | |