Star-Transformer: https://arxiv.org/pdf/1902.09113v2.pdf
| 67 | |
| 68 | |
| 69 | class StarEncoderLayer(nn.Module): |
| 70 | ''' Star-Transformer: https://arxiv.org/pdf/1902.09113v2.pdf ''' |
| 71 | |
| 72 | def __init__(self, d_model, n_head, d_k, d_v, dropout=0.1): |
| 73 | super(StarEncoderLayer, self).__init__() |
| 74 | self.slf_attn_satellite = MultiHeadAttention( |
| 75 | n_head, d_model, d_k, d_v, use_star=True, dropout=dropout) |
| 76 | self.slf_attn_relay = MultiHeadAttention( |
| 77 | n_head, d_model, d_k, d_v, use_star=True, dropout=dropout) |
| 78 | |
| 79 | def forward(self, h, e, s, non_pad_mask=None, slf_attn_mask=None): |
| 80 | # satellite node |
| 81 | batch_size, seq_len, d_model = h.size() |
| 82 | h_extand = torch.zeros(batch_size, seq_len+2, d_model, dtype=torch.float, device=h.device) |
| 83 | h_extand[:, 1:seq_len+1, :] = h # head and tail padding(not cycle) |
| 84 | s = s.reshape([batch_size, 1, d_model]) |
| 85 | s_expand = s.expand([batch_size, seq_len, d_model]) |
| 86 | context = torch.cat((h_extand[:, 0:seq_len, :], |
| 87 | h_extand[:, 1:seq_len+1, :], |
| 88 | h_extand[:, 2:seq_len+2, :], |
| 89 | e, |
| 90 | s_expand), |
| 91 | 2) |
| 92 | context = context.reshape([batch_size*seq_len, 5, d_model]) |
| 93 | h = h.reshape([batch_size*seq_len, 1, d_model]) |
| 94 | |
| 95 | h, _ = self.slf_attn_satellite( |
| 96 | h, context, context, mask=slf_attn_mask) |
| 97 | h = torch.squeeze(h, 1).reshape([batch_size, seq_len, d_model]) |
| 98 | if non_pad_mask is not None: |
| 99 | h *= non_pad_mask |
| 100 | |
| 101 | # virtual relay node |
| 102 | s_h = torch.cat((s, h), 1) |
| 103 | s, _ = self.slf_attn_relay( |
| 104 | s, s_h, s_h, mask=slf_attn_mask) |
| 105 | s = torch.squeeze(s, 1) |
| 106 | |
| 107 | return h, s |