PepPA / src /peppa /ptm.py
pranamanam's picture
Upload 97 files
98bde72 verified
Raw
History Blame Contribute Delete
3.31 kB
"""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)}