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