| 39 | } |
| 40 | |
| 41 | class TransformerEncoder(nn.Module): |
| 42 | |
| 43 | def __init__(self, encoder_layer, num_layers, |
| 44 | norm=None, weight_init_name="xavier_uniform"): |
| 45 | super().__init__() |
| 46 | self.layers = get_clones(encoder_layer, num_layers) |
| 47 | self.num_layers = num_layers |
| 48 | self.norm = norm |
| 49 | self._reset_parameters(weight_init_name) |
| 50 | |
| 51 | def _reset_parameters(self, weight_init_name): |
| 52 | func = WEIGHT_INIT_DICT[weight_init_name] |
| 53 | for p in self.parameters(): |
| 54 | if p.dim() > 1: |
| 55 | func(p) |
| 56 | |
| 57 | def forward(self, src, |
| 58 | mask: Optional[Tensor] = None, |
| 59 | src_key_padding_mask: Optional[Tensor] = None, |
| 60 | pos: Optional[Tensor] = None, |
| 61 | xyz: Optional [Tensor] = None, |
| 62 | transpose_swap: Optional[bool] = False, |
| 63 | return_attn_weights: Optional [bool] = False, |
| 64 | ): |
| 65 | attns = [] |
| 66 | if transpose_swap: |
| 67 | bs, c, h, w = src.shape |
| 68 | src = src.flatten(2).permute(2, 0, 1) |
| 69 | if pos is not None: |
| 70 | pos = pos.flatten(2).permute(2, 0, 1) |
| 71 | output = src |
| 72 | orig_mask = mask |
| 73 | if orig_mask is not None and isinstance(orig_mask, list): |
| 74 | assert len(orig_mask) == len(self.layers) |
| 75 | elif orig_mask is not None: |
| 76 | orig_mask = [mask for _ in range(len(self.layers))] |
| 77 | |
| 78 | for idx, layer in enumerate(self.layers): |
| 79 | if orig_mask is not None: |
| 80 | mask = orig_mask[idx] |
| 81 | # mask must be tiled to num_heads of the transformer |
| 82 | bsz, n, n = mask.shape |
| 83 | nhead = layer.nhead |
| 84 | mask = mask.unsqueeze(1) |
| 85 | mask = mask.repeat(1, nhead, 1, 1) |
| 86 | mask = mask.view(bsz * nhead, n, n) |
| 87 | if return_attn_weights: |
| 88 | output, attn = layer(output, src_mask=mask, |
| 89 | src_key_padding_mask=src_key_padding_mask, pos=pos, return_attn_weights=True) |
| 90 | attns.append(attn) |
| 91 | else: |
| 92 | output = layer(output, src_mask=mask, |
| 93 | src_key_padding_mask=src_key_padding_mask, pos=pos) |
| 94 | |
| 95 | if self.norm is not None: |
| 96 | output = self.norm(output) |
| 97 | |
| 98 | if transpose_swap: |