| """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 |
| |
| 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())) |
|
|