| 70 | return msa |
| 71 | |
| 72 | class PairStr2Pair(nn.Module): |
| 73 | def __init__(self, d_pair=128, n_head=4, d_hidden=32, d_rbf=36, p_drop=0.15): |
| 74 | super(PairStr2Pair, self).__init__() |
| 75 | |
| 76 | self.emb_rbf = nn.Linear(d_rbf, d_hidden) |
| 77 | self.proj_rbf = nn.Linear(d_hidden, d_pair) |
| 78 | |
| 79 | self.drop_row = Dropout(broadcast_dim=1, p_drop=p_drop) |
| 80 | self.drop_col = Dropout(broadcast_dim=2, p_drop=p_drop) |
| 81 | |
| 82 | self.row_attn = BiasedAxialAttention(d_pair, d_pair, n_head, d_hidden, p_drop=p_drop, is_row=True) |
| 83 | self.col_attn = BiasedAxialAttention(d_pair, d_pair, n_head, d_hidden, p_drop=p_drop, is_row=False) |
| 84 | |
| 85 | self.ff = FeedForwardLayer(d_pair, 2) |
| 86 | |
| 87 | self.reset_parameter() |
| 88 | |
| 89 | def reset_parameter(self): |
| 90 | nn.init.kaiming_normal_(self.emb_rbf.weight, nonlinearity='relu') |
| 91 | nn.init.zeros_(self.emb_rbf.bias) |
| 92 | |
| 93 | self.proj_rbf = init_lecun_normal(self.proj_rbf) |
| 94 | nn.init.zeros_(self.proj_rbf.bias) |
| 95 | |
| 96 | def forward(self, pair, rbf_feat): |
| 97 | B, L = pair.shape[:2] |
| 98 | |
| 99 | rbf_feat = self.proj_rbf(F.relu_(self.emb_rbf(rbf_feat))) |
| 100 | |
| 101 | pair = pair + self.drop_row(self.row_attn(pair, rbf_feat)) |
| 102 | pair = pair + self.drop_col(self.col_attn(pair, rbf_feat)) |
| 103 | pair = pair + self.ff(pair) |
| 104 | return pair |
| 105 | |
| 106 | class MSA2Pair(nn.Module): |
| 107 | def __init__(self, d_msa=256, d_pair=128, d_hidden=32, p_drop=0.15): |