import argparse import pandas as pd import numpy as np from collections import defaultdict from itertools import chain from pathlib import Path import zarr from zarr.storage import LocalStore # Load reference data seq_exist = pd.read_csv('/app/MANE_all_transcipts.csv') def get_round_2(threshold, workdir, round, drug_csv: Path): workdir = Path(workdir) fasta_name = "patient_0" drug_name = drug_csv.stem name_out = f"{drug_name}_{fasta_name}" # Load all zarr score stores produced by screen.py. results_zarr = list(workdir.glob('*_results.zarr')) if not results_zarr: raise FileNotFoundError("No results Zarr stores found") print("Loading Zarr results...") # Lazy zarr handles per file; we read the (proteins, compounds) score matrix # in protein chunks rather than loading it fully — keeps peak memory bounded # regardless of compound count. zarr_groups = [] compound_ids_per = [] protein_ids_per = [] for z_path in results_zarr: store = LocalStore(z_path, read_only=True) group = zarr.open(store, mode="r") zarr_groups.append(group) # numpy 2.4+ returns StringDType for these arrays; convert via tolist() # because direct .astype(str) raises "cannot cast StringDType to StrDType". compound_ids_per.append(np.array(group["compound_ids"][:].tolist(), dtype=str)) protein_ids_per.append(np.array(group["protein_ids"][:].tolist(), dtype=str)) # All zarrs share the same compound axis (drug + same metabolites). all_compound_ids = compound_ids_per[0] for cids in compound_ids_per[1:]: if cids.shape != all_compound_ids.shape or not np.array_equal(cids, all_compound_ids): raise ValueError("All *_results.zarr stores must share the same compound_ids axis") all_protein_ids = np.concatenate(protein_ids_per) if len(protein_ids_per) > 1 else protein_ids_per[0] n_proteins = sum(g["scores"].shape[0] for g in zarr_groups) n_compounds = all_compound_ids.shape[0] # ---- Compound-axis indices (tiny) ---- drug_mask_cmpd = np.array([s.startswith("drug") for s in all_compound_ids], dtype=bool) if not drug_mask_cmpd.any(): raise ValueError("No 'drug_*' compounds present in compound_ids") drug_cols_idx = np.where(drug_mask_cmpd)[0] drug0_search = np.where(all_compound_ids == 'drug_0')[0] if drug0_search.size == 0: raise ValueError("'drug_0' not found in compound_ids") drug0_idx = int(drug0_search[0]) # ---- Stream the score matrix in protein chunks, writing significant_interactions # incrementally and accumulating per-protein summary stats. Avoids both the # full score matrix load and the giant in-memory above-threshold DataFrame. ---- # Per chunk peak ≈ CHUNK_PROTEINS * n_compounds * 4 bytes; for 1000×38k ≈ 150 MB. CHUNK_PROTEINS = 1000 drug_position = np.zeros(n_proteins, dtype=np.int64) drug_0_col_full = np.empty(n_proteins, dtype=np.float32) prot_count = defaultdict(int) # protein_id -> count of above-threshold compounds prot_sum = defaultdict(float) # protein_id -> sum of above-threshold scores # Open significant_interactions.tsv before the loop (round 2 only); write rows per chunk. sig_path = workdir / f'{name_out}_significant_interactions.tsv' sig_writer = None smile_dict = None if round == 2: test_smiles = pd.read_csv(workdir / 'smiles.smi', sep='\t', header=None) smile_dict = dict(zip(test_smiles[1], test_smiles[0])) sig_writer = open(sig_path, 'w') sig_writer.write("drug/metabolite\tsmile\ttranscipt\tconplex_score\n") cursor = 0 for group in zarr_groups: scores_arr = group["scores"] n_this = scores_arr.shape[0] for start in range(0, n_this, CHUNK_PROTEINS): end = min(start + CHUNK_PROTEINS, n_this) chunk = np.asarray(scores_arr[start:end]) # (chunk_size, n_compounds) global_start = cursor + start global_end = cursor + end chunk_size = end - start # drug score columns + max-per-protein for this chunk chunk_max_drug = chunk[:, drug_cols_idx].max(axis=1) # drug_position[i] = # compounds with score strictly > max_drug_score_i. # Matches production's stable-sort + first-drug-position convention. drug_position[global_start:global_end] = (chunk > chunk_max_drug[:, None]).sum(axis=1) # drug_0 column drug_0_col_full[global_start:global_end] = chunk[:, drug0_idx].astype(np.float32, copy=False) # Above-threshold rows for this chunk pi_local, ci = np.where(chunk > threshold) if pi_local.size: scores_vec = chunk[pi_local, ci].astype(np.float32, copy=False) # Per-protein aggregates accumulated via bincount on local indices. counts_local = np.bincount(pi_local, minlength=chunk_size) sums_local = np.bincount(pi_local, weights=scores_vec.astype(np.float64), minlength=chunk_size) for pi in np.flatnonzero(counts_local): pid = all_protein_ids[global_start + pi] prot_count[pid] += int(counts_local[pi]) prot_sum[pid] += float(sums_local[pi]) # Write significant_interactions rows for this chunk directly to the file. if sig_writer is not None: compound_ids_vec = all_compound_ids[ci] protein_ids_vec = all_protein_ids[global_start + pi_local] smiles_vec = np.array( [smile_dict.get(c, "UNKNOWN") for c in compound_ids_vec] ) pd.DataFrame({ "drug/metabolite": compound_ids_vec, "smile": smiles_vec, "transcipt": protein_ids_vec, "conplex_score": scores_vec, }).to_csv(sig_writer, sep='\t', index=False, header=False) cursor += n_this if sig_writer is not None: sig_writer.close() drug_0_score = pd.DataFrame({ "Drug": "drug_0", "Transcript": all_protein_ids, "Score": drug_0_col_full, }) if round == 2: drug_0_score.to_csv(workdir / f'{name_out}_drug_scores.tsv', sep='\t', index=False) # ---- Per-protein summary (only proteins with any above-threshold compound) ---- drug_pos_map = dict(zip(all_protein_ids, drug_position)) if round == 1 and not prot_count: # Equivalent to old "transcipts_summary.empty" early-out for round 1. pd.DataFrame().to_csv(workdir / 'round_1.csv', index=False) with open(workdir / 'round_1.fasta', 'w') as f: f.write("all_data_is_filtered_out\n") return 'STOP NO TRANSCIPTS ABOVE THRESHOLD' if prot_count: proteins_with_above = list(prot_count.keys()) # mean dtype kept as float32 to match the precision of the previous # groupby-on-float32 path (CSV output is bit-identical that way). means = np.array( [prot_sum[p] / prot_count[p] for p in proteins_with_above], dtype=np.float32, ) summary = pd.DataFrame({ 'transcipt_name': proteins_with_above, 'number_of_iteracting_compounds': [prot_count[p] for p in proteins_with_above], 'mean_binding_above_threshold': means, }) else: summary = pd.DataFrame(columns=[ 'transcipt_name', 'number_of_iteracting_compounds', 'mean_binding_above_threshold' ]) summary['drug_position'] = summary['transcipt_name'].map(drug_pos_map) path_out = workdir / f'round_{round}.csv' if len(summary) == 0: with path_out.open("w") as f: f.write("NO TRANSCIPTS ABOVE THRESHOLD") return 'STOP NO TRANSCIPTS ABOVE THRESHOLD' summary['if_drug_above_threshold'] = ( summary['number_of_iteracting_compounds'] > summary['drug_position'] ) # MANE symbol lookup; strip mutated-protein "_2" suffix. clean_transcript_names = [str(t).split('_')[0] for t in summary['transcipt_name']] seq_lookup = seq_exist.set_index('transcipt')['symbol'] summary['protein_name'] = [ seq_lookup.get(t, "") for t in clean_transcript_names ] summary.to_csv(path_out, index=False) # ---- Round 1 also writes the round-2 FASTA ---- if round == 1: name_2_filtered = list(summary['protein_name']) transcipts_2 = list(chain.from_iterable( list(seq_exist[seq_exist['symbol'] == sym]['transcipt']) for sym in name_2_filtered )) transcipts_2 = list(np.unique(transcipts_2)) existing = set(summary['transcipt_name']) transcipts_2 = [t for t in transcipts_2 if t not in existing] fasta_new = [] for t in transcipts_2: row = seq_exist[seq_exist['transcipt'] == t] if not row.empty: fasta_new.append(f">{t}") fasta_new.append(row['seq'].iloc[0]) with open(workdir / 'round_2.fasta', 'w') as f: for line in fasta_new: f.write(f"{line}\n") print(f"Round {round} processing complete.") if __name__ == "__main__": parser = argparse.ArgumentParser(description="Process and screen fragments.") parser.add_argument("--threshold", required=True, type=float, help="Threshold for 1. round of conplex scores") parser.add_argument("--workdir", required=True, type=Path, help="workdir") parser.add_argument("--round", required=True, type=int, help="Round 1 or 2") parser.add_argument("--drug-csv", required=True, type=Path, help="csv file with drug smiles") args = parser.parse_args() get_round_2(args.threshold, args.workdir, args.round, args.drug_csv)