(self, d_pair=128, n_head=4, d_hidden=32, d_rbf=36, p_drop=0.15)
| 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') |
nothing calls this directly
no test coverage detected