File size: 3,305 Bytes
98bde72
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
"""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)}