ReMEDi / scripts /fetch_sciplex_subset.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw
History Blame Contribute Delete
6.38 kB
"""Fetch a bounded sci-Plex3 prototype by HTTP ranges, with exact-name Tahoe SMILES.
The source is the corrected scPerturb object. The selected cohort is intentionally
small and is an engineering prototype, not the external ICML benchmark.
"""
import argparse
from remedi.download import RangeReader
from pathlib import Path
import json
import hashlib
import time
import numpy as np
import pandas as pd
import requests
import fsspec
import h5py
import anndata as ad
import pyarrow.parquet as pq
from scipy import sparse
URL = "https://zenodo.org/records/13350497/files/SrivatsanTrapnell2020_sciplex3.h5ad?download=1"
DRUG_URL = "https://huggingface.co/datasets/tahoebio/Tahoe-100M/resolve/main/metadata/drug_metadata.parquet"
def read_column(node):
if isinstance(node, h5py.Group):
categories = read_column(node["categories"])
codes = node["codes"][:]
out = categories[np.maximum(codes, 0)].astype(object)
out[codes < 0] = None
return out
values = node[:]
return np.array([v.decode() if isinstance(v, bytes) else v for v in values]) if values.dtype.kind in "SO" else values
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--output", required=True)
parser.add_argument("--source", default=URL, help="Local corrected h5ad or source URL")
parser.add_argument("--cells-per-condition", type=int, default=8)
parser.add_argument("--cell-line", default="A549")
parser.add_argument("--max-molecules", type=int, default=40)
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--metadata-cache", help="Optional locally generated metadata cache directory")
args = parser.parse_args()
output = Path(args.output); output.mkdir(parents=True, exist_ok=True)
cache = Path(args.metadata_cache) if args.metadata_cache else output
drug_file = cache/"drug_metadata.parquet"
if not drug_file.exists():
response = requests.get(DRUG_URL, timeout=60); response.raise_for_status()
drug_file.parent.mkdir(parents=True, exist_ok=True); drug_file.write_bytes(response.content)
# Arrow conversion avoids dependence on a pandas-specific parquet extension registry.
drugs = pd.DataFrame(pq.read_table(drug_file).to_pylist())
lookup = {s.strip().casefold(): i for i, s in enumerate(drugs.drug) if isinstance(s, str)}
rng = np.random.default_rng(args.seed)
start = time.perf_counter()
handle = RangeReader(args.source, output/'range_cache', 2526631614) if args.source.startswith('https://') else open(args.source, 'rb')
with handle:
with h5py.File(handle, "r") as h:
names = ["cell_line","dose_value","perturbation","replicate","plate","well","time"]
obs = pd.DataFrame({name: read_column(h["obs"][name]) for name in names})
print(f"Loaded metadata for {len(obs)} cells", flush=True)
matches = sorted(s for s in obs.perturbation.dropna().unique() if s.strip().casefold() in lookup)[:args.max_molecules]
mappings = [{"drug": s, "smiles": drugs.iloc[lookup[s.strip().casefold()]].canonical_smiles,
"tahoe_name": drugs.iloc[lookup[s.strip().casefold()]].drug,
"matching_rule": "exact case-folded name after outer whitespace removal"} for s in matches]
pd.DataFrame(mappings).to_csv(output/"structures.csv", index=False)
eligible = obs[(obs.cell_line == args.cell_line) & (obs.time == 24.)]
treated = eligible[eligible.perturbation.isin(matches) & eligible.dose_value.isin([100.,1000.])]
selected = []
for _, group in treated.groupby(["perturbation","dose_value","replicate","plate"], sort=True):
if len(group) >= args.cells_per_condition:
selected.extend(rng.choice(group.index, args.cells_per_condition, replace=False))
plates = set(obs.loc[selected, "plate"])
controls = eligible[(eligible.perturbation == "control") & eligible.plate.isin(plates)]
for _, group in controls.groupby(["replicate","plate"], sort=True):
selected.extend(rng.choice(group.index, min(args.cells_per_condition*2,len(group)), replace=False))
selected = np.sort(np.unique(selected))
x = h["X"]
if x.attrs.get("encoding-type") != "csr_matrix": raise ValueError("Expected source CSR matrix")
pointers = x["indptr"][:]
values, columns, pointer = [], [], [0]
for j, index in enumerate(selected):
begin, end = int(pointers[index]), int(pointers[index+1])
values.append(x["data"][begin:end]); columns.append(x["indices"][begin:end])
pointer.append(pointer[-1]+end-begin)
if j % 100 == 0: print(f"Read {j}/{len(selected)} selected cell rows", flush=True)
matrix = sparse.csr_matrix((np.concatenate(values),np.concatenate(columns),np.asarray(pointer)),
shape=(len(selected),int(x.attrs["shape"][1])))
var_index = h["var"].attrs.get("_index", "_index")
if isinstance(var_index, bytes): var_index=var_index.decode()
gene_names = read_column(h["var"][var_index]).astype(str)
subset_obs = obs.iloc[selected].copy()
subset_obs["source_row"] = selected
subset_obs.index = [f"sciplex3_row_{i}" for i in selected]
result = ad.AnnData(matrix, obs=subset_obs, var=pd.DataFrame(index=gene_names))
result.write_h5ad(output/"sciplex_prototype.h5ad", compression="gzip")
manifest = {"source_url":URL,"source_record":"scPerturb corrected v1.4 / Zenodo 13350497",
"source_md5":"c9e70629505d98c7ca1a837f62b14e89", "full_source_checksum_recomputed":False,
"smiles_source":DRUG_URL,"molecules":len(matches),"cells":len(selected),
"gene_columns":matrix.shape[1],"nonzero_counts":matrix.nnz,"cell_line":args.cell_line,
"exposure_hours":24,"dose_um":[.1,1.],"seed":args.seed,
"prototype_only":True,"elapsed_seconds":time.perf_counter()-start,
"selected_rows_sha256":hashlib.sha256(selected.tobytes()).hexdigest()}
(output/"source_manifest.json").write_text(json.dumps(manifest, indent=2)+"\n")
print(json.dumps(manifest,indent=2))
if __name__ == "__main__": main()