File size: 2,578 Bytes
d766458
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
69
70
71
72
73
74
75
76
77
78
import jax
import jax.numpy as jnp
import numpy as np

def get_stats(X, X_weight=None, labels=None, add_f_ij=True, add_mf_ij=False, add_c=False):
  '''compute f_i/f_ij/f_ijk given msa '''
  n = None
  if X_weight is None:
    Xn = Xs = X
  else:
    Xn, Xs = X*X_weight[:,n,n], X*jnp.sqrt(X_weight[:,n,n])
  f_i = Xn.sum(0)
  o = {"f_i": f_i / f_i.sum(1,keepdims=True)}
  
  if add_f_ij:
    f_ij = jnp.tensordot(Xs,Xs,[0,0])
    o["f_ij"] = f_ij / f_ij.sum((1,3),keepdims=True)
    if add_c: o["c_ij"] = o["f_ij"] - o["f_i"][:,:,n,n] * o["f_i"][n,n,:,:]

  if labels is not None:
    # compute mixture stats
    if jnp.issubdtype(labels, jnp.integer):
      labels = jax.nn.one_hot(labels,labels.max()+1)
    mf_i = jnp.einsum("nc,nia->cia", labels, Xn)
    o["mf_i"] = mf_i/mf_i.sum((0,2),keepdims=True)
    if add_mf_ij:
      mf_ij = jnp.einsum("nc,nia,njb->ciajb", labels, Xs, Xs)
      o["mf_ij"] = mf_ij/mf_ij.sum((0,2,4),keepdims=True)   
      if add_c: o["mc_ij"] = o["mf_ij"] - o["mf_i"][:,:,:,n,n] * o["mf_i"][:,n,n,:,:]
  return o
    
def get_r(a,b):
  a = jnp.array(a).flatten()
  b = jnp.array(b).flatten()
  return jnp.corrcoef(a,b)[0,1]

def inv_cov(X, X_weight=None):
  X = jnp.asarray(X)
  N,L,A = X.shape
  if X_weight is None:
    num_points = N
  else:
    X_weight = jnp.asarray(X_weight)
    num_points = X_weight.sum()
  c = get_stats(X, X_weight, add_mf_ij=True, add_c=True)["c_ij"]
  c = c.reshape(L*A,L*A)
  shrink = 4.5/jnp.sqrt(num_points) * jnp.eye(c.shape[0])
  ic = jnp.linalg.inv(c + shrink)
  return ic.reshape(L,A,L,A)

def get_mtx(W):
  W = jnp.asarray(W)
  # l2norm of 20x20 matrices (note: we ignore gaps)
  raw = jnp.sqrt(jnp.sum(np.square(W[:,1:,:,1:]),(1,3)))
  raw = raw.at[jnp.diag_indices_from(raw)].set(0)

  # apc (average product correction)
  ap = raw.sum(0,keepdims=True) * raw.sum(1,keepdims=True) / raw.sum()
  apc = raw - ap
  apc = apc.at[jnp.diag_indices_from(apc)].set(0)
  return raw, apc

def con_auc(true, pred, mask=None):
  '''compute agreement between predicted and measured contact map'''
  true = jnp.asarray(true)
  pred = jnp.asarray(pred)
  if mask is not None:
    mask = jnp.asarray(mask)
    idx = mask.sum(-1) > 0
    true = true[idx,:][:,idx]
    pred = pred[idx,:][:,idx]
  eval_idx = jnp.triu_indices_from(true, 6)
  pred_, true_ = pred[eval_idx], true[eval_idx] 
  L = (jnp.linspace(0.1,1.0,10)*len(true)).astype(jnp.int32)
  sort_idx = jnp.argsort(pred_)[::-1]
  return jnp.asarray([true_[sort_idx[:l]].mean() for l in L])