| 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): |
| 108 | super(MSA2Pair, self).__init__() |
| 109 | self.norm = nn.LayerNorm(d_msa) |
| 110 | self.proj_left = nn.Linear(d_msa, d_hidden) |
| 111 | self.proj_right = nn.Linear(d_msa, d_hidden) |
| 112 | self.proj_out = nn.Linear(d_hidden*d_hidden, d_pair) |
| 113 | |
| 114 | self.reset_parameter() |
| 115 | |
| 116 | def reset_parameter(self): |
| 117 | # normal initialization |
| 118 | self.proj_left = init_lecun_normal(self.proj_left) |
| 119 | self.proj_right = init_lecun_normal(self.proj_right) |
| 120 | nn.init.zeros_(self.proj_left.bias) |
| 121 | nn.init.zeros_(self.proj_right.bias) |
| 122 | |
| 123 | # zero initialize output |
| 124 | nn.init.zeros_(self.proj_out.weight) |
| 125 | nn.init.zeros_(self.proj_out.bias) |
| 126 | |
| 127 | def forward(self, msa, pair): |
| 128 | B, N, L = msa.shape[:3] |
| 129 | msa = self.norm(msa) |
| 130 | left = self.proj_left(msa) |
| 131 | right = self.proj_right(msa) |
| 132 | right = right / float(N) |
| 133 | out = einsum('bsli,bsmj->blmij', left, right).reshape(B, L, L, -1) |
| 134 | out = self.proj_out(out) |
| 135 | |
| 136 | pair = pair + out |
| 137 | |
| 138 | return pair |
| 139 | |
| 140 | class SCPred(nn.Module): |
| 141 | def __init__(self, d_msa=256, d_state=32, d_hidden=128, p_drop=0.15): |