| """Supervised PTM interaction head on externally computed frozen embeddings. |
| |
| Unknown interactions carry NaN labels and never enter classification loss. |
| Paired indices refer to the same peptide and matched target chemistry. |
| """ |
| from dataclasses import dataclass |
| import numpy as np |
| from scipy.optimize import minimize |
| from scipy.special import expit |
|
|
| @dataclass |
| class PTMHead: |
| weight: np.ndarray |
| bias: float |
| mean_b: np.ndarray |
| std_b: np.ndarray |
| mean_t: np.ndarray |
| std_t: np.ndarray |
|
|
| def logits(self, binder, target): |
| b=(np.asarray(binder)-self.mean_b)/self.std_b |
| t=(np.asarray(target)-self.mean_t)/self.std_t |
| return np.einsum('ni,ij,nj->n',b,self.weight,t)+self.bias |
|
|
| def save(self,path): |
| np.savez(path,weight=self.weight,bias=self.bias,mean_b=self.mean_b, |
| std_b=self.std_b,mean_t=self.mean_t,std_t=self.std_t) |
| @classmethod |
| def load(cls,path): |
| with np.load(path,allow_pickle=False) as x: |
| return cls(**{k:x[k] for k in x.files}) |
|
|
|
|
| def loss_gradient(theta,b,t,labels,pairs,pair_weight=1.,margin=1.,l2=1e-3): |
| """BCE on observed labels + pairwise hinge + Frobenius regularization.""" |
| w=theta[:-1].reshape(b.shape[1],t.shape[1]); z=np.einsum('ni,ij,nj->n',b,w,t)+theta[-1] |
| mask=np.isfinite(labels); dz=np.zeros(len(z));loss=0. |
| if mask.any(): |
| y=labels[mask] |
| if not np.isin(y,[0,1]).all():raise ValueError('observed labels must be binary') |
| loss=float(np.mean(np.logaddexp(0,z[mask])-y*z[mask])) |
| dz[mask]=(expit(z[mask])-y)/mask.sum() |
| if len(pairs): |
| pos,neg=np.asarray(pairs,dtype=int).T |
| violation=margin-z[pos]+z[neg];active=violation>0 |
| loss+=pair_weight*np.maximum(violation,0).mean() |
| np.add.at(dz,pos[active],-pair_weight/len(pairs)) |
| np.add.at(dz,neg[active],pair_weight/len(pairs)) |
| loss+=l2*np.sum(w*w) |
| gw=np.einsum('n,ni,nj->ij',dz,b,t)+2*l2*w |
| return loss,np.r_[gw.ravel(),dz.sum()] |
|
|
|
|
| def fit(binder,target,labels,pairs=(),pair_weight=1.,margin=1.,l2=1e-3,maxiter=300): |
| b,t=np.asarray(binder,dtype=float),np.asarray(target,dtype=float);y=np.asarray(labels,dtype=float) |
| if b.ndim!=2 or t.ndim!=2 or len(b)!=len(t) or y.shape!=(len(b),): |
| raise ValueError('expected aligned N x d embedding matrices and N labels') |
| if not np.isfinite(b).all() or not np.isfinite(t).all():raise ValueError('nonfinite embedding') |
| if not np.isfinite(y).any() and not len(pairs):raise ValueError('no supervised observations') |
| if len(pairs) and (np.min(pairs)<0 or np.max(pairs)>=len(b)):raise ValueError('pair index outside training rows') |
| mb,sb=b.mean(0),np.maximum(b.std(0),1e-6);mt,st=t.mean(0),np.maximum(t.std(0),1e-6) |
| bn,tn=(b-mb)/sb,(t-mt)/st |
| result=minimize(loss_gradient,np.zeros(b.shape[1]*t.shape[1]+1),args=(bn,tn,y,pairs,pair_weight,margin,l2), |
| jac=True,method='L-BFGS-B',options={'maxiter':maxiter,'ftol':1e-10}) |
| if not result.success:raise RuntimeError('PTM head optimization failed: '+result.message) |
| head=PTMHead(result.x[:-1].reshape(b.shape[1],t.shape[1]),float(result.x[-1]),mb,sb,mt,st) |
| return head,{'loss':float(result.fun),'iterations':int(result.nit),'observed_labels':int(np.isfinite(y).sum()),'paired_examples':len(pairs)} |
|
|