ReMEDi / scripts /evaluate_forward.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw
History Blame Contribute Delete
2.35 kB
"""Compare response-head predictions with zero, training mean, and chemical neighbors."""
import argparse
from pathlib import Path
import joblib
import numpy as np
import pandas as pd
from remedi.io import ResponseData
from remedi.models import predict,check_feature_space
from remedi.chemistry import fingerprints
p=argparse.ArgumentParser()
for name in ['data','model','output']:p.add_argument('--'+name,required=True)
p.add_argument('--split',default='test');p.add_argument('--embeddings')
a=p.parse_args();d=ResponseData.load(a.data);m=joblib.load(Path(a.model)/'model.joblib');check_feature_space(m,d)
tr=np.flatnonzero(d.obs.split.to_numpy()=='train');te=np.flatnonzero(d.obs.split.to_numpy()==a.split)
if not len(tr) or not len(te):raise ValueError('Training and evaluation conditions are required')
o=d.obs.iloc[te];y=d.response[te]
preds={'zero':np.zeros_like(y),'training-mean':np.repeat(d.response[tr].mean(0)[None,:],len(te),axis=0),
'response-head':predict(m,o.smiles,o.dose_um,d.control[te],a.embeddings).mean(0)}
fp=fingerprints(d.obs.smiles);nearest=[]
for i in te:
candidates=tr[d.obs.iloc[tr].context.to_numpy()==d.obs.iloc[i].context]
if not len(candidates):candidates=tr
overlap=fp[candidates]@fp[i];union=fp[candidates].sum(1)+fp[i].sum()-overlap
similarity=np.divide(overlap,union,out=np.zeros_like(overlap),where=union>0)
# Select chemical neighbors first and the closest observed log dose within ties.
top=candidates[np.isclose(similarity,similarity.max())]
j=top[np.argmin(abs(np.log(d.obs.iloc[top].dose_um.to_numpy()/d.obs.iloc[i].dose_um)))]
nearest.append(d.response[j])
preds['chemical-neighbor']=np.asarray(nearest)
rows=[]
for name,values in preds.items():
for k,i in enumerate(te):
truth=y[k];estimate=values[k]
corr=float(np.corrcoef(truth,estimate)[0,1]) if truth.std()>0 and estimate.std()>0 else np.nan
rows.append({'condition':int(i),'molecule_id':d.obs.iloc[i].molecule_id,'dose_um':d.obs.iloc[i].dose_um,
'method':name,'mse':float(np.mean((estimate-truth)**2)),'pearson':corr})
out=Path(a.output);out.mkdir(parents=True,exist_ok=True)
r=pd.DataFrame(rows);r.to_csv(out/'forward_conditions.csv',index=False)
r.groupby(['method','molecule_id'])[['mse','pearson']].mean().groupby('method').mean().to_csv(out/'forward_summary.csv')