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