| 89 | return emb_out |
| 90 | |
| 91 | class PositionalEncoding2D(nn.Module): |
| 92 | # Add relative positional encoding to pair features |
| 93 | def __init__(self, d_model, minpos=-32, maxpos=32, p_drop=0.1): |
| 94 | super(PositionalEncoding2D, self).__init__() |
| 95 | self.minpos = minpos |
| 96 | self.maxpos = maxpos |
| 97 | self.nbin = abs(minpos)+maxpos+1 |
| 98 | self.emb = nn.Embedding(self.nbin, d_model) |
| 99 | self.drop = nn.Dropout(p_drop) |
| 100 | |
| 101 | def forward(self, x, idx): |
| 102 | bins = torch.arange(self.minpos, self.maxpos, device=x.device) |
| 103 | seqsep = idx[:,None,:] - idx[:,:,None] # (B, L, L) |
| 104 | # |
| 105 | ib = torch.bucketize(seqsep, bins).long() # (B, L, L) |
| 106 | emb = self.emb(ib) #(B, L, L, d_model) |
| 107 | x = x + emb # add relative positional encoding |
| 108 | return self.drop(x) |
| 109 | |
| 110 | class MSA_emb(nn.Module): |
| 111 | # Get initial seed MSA embedding |