File size: 4,289 Bytes
3f98d52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
"""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()))