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