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