ReMEDi / scripts /prepare_tahoe.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw
History Blame Contribute Delete
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()))