File size: 3,184 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
79
80
81
82
83
84
85
86
87
88
89
import jax
import jax.numpy as jnp
import numpy

from colabdesign.seq.kmeans import kmeans
from colabdesign.seq.stats import get_stats, get_eff

# LEARN SEQUENCES
# "parameter-free" model, where we learn msa to match the statistics.
# We can take kmeans to the "next" level and directly optimize sequences to match desired stats.

class LEARN_MSA:
  def __init__(self, X, X_weight=None, samples=None,               

               mode="tied", k=1,

               seed=0, learning_rate=1e-3):
    
    assert mode in ["tied","full"]
    assert k > 0

    key = jax.random.PRNGKey(seed)
    self.k = k

    # collect X stats 
    N,L,A = X.shape
    if samples is None: samples = N
    X = jnp.asarray(X)
    X_weight = get_eff(X) if X_weight is None else jnp.asarray(X_weight)

    # run kmeans
    self.kms = kmeans(X, X_weight, k=self.k)
    stats_args = dict(add_f_ij=True, add_mf_ij=(mode=="full"), add_c=True)
    self.X_stats = get_stats(X, X_weight, labels=jax.nn.one_hot(self.kms["labels"],self.k), **stats_args)

    if samples == N:
      self.Y_labels = self.kms["labels"]
    else:
      # sample labels
      key,key_ = jax.random.split(key)
      self.Y_labels = jnp.sort(jax.random.choice(key_, jnp.arange(k), shape=(samples,), p=self.kms["cat"]))

    key, key_ = jax.random.split(key)
    Neff = X_weight.sum()
    Y_logits = jnp.log(self.kms["means"] * Neff + 0.01 * jnp.log(Neff))[self.Y_labels]
    Y = jax.nn.softmax(Y_logits + jax.random.gumbel(key_,(samples,L,A)))

    # setup the model
    def model(params, X_stats):
      # categorical reparameterization of Y
      Y_hard = jax.nn.one_hot(params["Y"].argmax(-1),A)
      Y = jax.lax.stop_gradient(Y_hard - params["Y"]) + params["Y"]

      # collect Y stats
      Y_stats = get_stats(Y, labels=jax.nn.one_hot(self.Y_labels, self.k), **stats_args)
      
      # define loss function
      i,ij = ("f_i","c_ij") if k == 1 else ("mf_i",("c_ij" if mode == "tied" else "mc_ij"))
      loss_i = jnp.square(X_stats[i] - Y_stats[i]).sum((-1,-2))
      loss_ij = jnp.square(X_stats[ij] - Y_stats[ij]).sum((-1,-2,-3)).mean(-1)
      
      if self.k > 1:
        loss_i = (loss_i * self.kms["cat"]).sum()
        if mode == "full":
          loss_ij = (loss_ij * self.kms["cat"]).sum()
      
      loss = loss_i + loss_ij

      aux = {"r":get_r(X_stats["c_ij"], Y_stats["c_ij"])}
      return loss, aux

    # setup optimizer
    self.n = 0
    init_fun, self.update_fun, self.get_params = adam(learning_rate)
    self.state = init_fun({"Y":Y})
    self.grad = jax.jit(jax.value_and_grad(model, has_aux=True))

  def get_msa(self):
    Y = np.array(self.get_params(self.state)["Y"])
    return {"kms":self.kms,
            "sampled_msa":Y.argmax(-1),
            "sampled_labels":self.Y_labels}
      
  def fit(self, steps=100, verbose=True):
    '''train model'''
    for n in range(steps):
      (loss, aux), grad = self.grad(self.get_params(self.state), self.X_stats)
      self.state = self.update_fun(self.n, grad, self.state)
      self.n += 1
      if (n+1) % (steps // 10) == 0:
        print(self.n, loss, aux["r"])