| 68 | |
| 69 | |
| 70 | class CrossAttentionLayer(nn.Module): |
| 71 | |
| 72 | def __init__(self, d_model, nhead, dropout=0.0, |
| 73 | activation="relu", normalize_before=False): |
| 74 | super().__init__() |
| 75 | self.multihead_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout) |
| 76 | |
| 77 | self.norm = nn.LayerNorm(d_model) |
| 78 | self.dropout = nn.Dropout(dropout) |
| 79 | |
| 80 | self.activation = _get_activation_fn(activation) |
| 81 | self.normalize_before = normalize_before |
| 82 | |
| 83 | self._reset_parameters() |
| 84 | |
| 85 | def _reset_parameters(self): |
| 86 | for p in self.parameters(): |
| 87 | if p.dim() > 1: |
| 88 | nn.init.xavier_uniform_(p) |
| 89 | |
| 90 | def with_pos_embed(self, tensor, pos: Optional[Tensor]): |
| 91 | return tensor if pos is None else tensor + pos |
| 92 | |
| 93 | def forward_post(self, tgt, memory, |
| 94 | memory_mask: Optional[Tensor] = None, |
| 95 | memory_key_padding_mask: Optional[Tensor] = None, |
| 96 | pos: Optional[Tensor] = None, |
| 97 | query_pos: Optional[Tensor] = None): |
| 98 | tgt2, avg_attn = self.multihead_attn(query=self.with_pos_embed(tgt, query_pos), |
| 99 | key=self.with_pos_embed(memory, pos), |
| 100 | value=memory, attn_mask=memory_mask, |
| 101 | key_padding_mask=memory_key_padding_mask) |
| 102 | tgt = tgt + self.dropout(tgt2) |
| 103 | tgt = self.norm(tgt) |
| 104 | return tgt, avg_attn |
| 105 | |
| 106 | def forward_pre(self, tgt, memory, |
| 107 | memory_mask: Optional[Tensor] = None, |
| 108 | memory_key_padding_mask: Optional[Tensor] = None, |
| 109 | pos: Optional[Tensor] = None, |
| 110 | query_pos: Optional[Tensor] = None): |
| 111 | tgt2 = self.norm(tgt) |
| 112 | tgt2, avg_attn = self.multihead_attn(query=self.with_pos_embed(tgt2, query_pos), |
| 113 | key=self.with_pos_embed(memory, pos), |
| 114 | value=memory, attn_mask=memory_mask, |
| 115 | key_padding_mask=memory_key_padding_mask) |
| 116 | tgt = tgt + self.dropout(tgt2) |
| 117 | |
| 118 | return tgt, avg_attn |
| 119 | |
| 120 | def forward(self, tgt, memory, |
| 121 | memory_mask: Optional[Tensor] = None, |
| 122 | memory_key_padding_mask: Optional[Tensor] = None, |
| 123 | pos: Optional[Tensor] = None, |
| 124 | query_pos: Optional[Tensor] = None): |
| 125 | if self.normalize_before: |
| 126 | return self.forward_pre(tgt, memory, memory_mask, |
| 127 | memory_key_padding_mask, pos, query_pos) |
nothing calls this directly
no outgoing calls
no test coverage detected