(self,
in_channels,
out_channels,
nhead=8,
num_layers=4,
dim_feedforward=2048,
dropout=0.1,
max_length=25,
ignore_index=100,
pretraining=False,
detach=True)
| 13 | class BUSDecoder(nn.Module): |
| 14 | |
| 15 | def __init__(self, |
| 16 | in_channels, |
| 17 | out_channels, |
| 18 | nhead=8, |
| 19 | num_layers=4, |
| 20 | dim_feedforward=2048, |
| 21 | dropout=0.1, |
| 22 | max_length=25, |
| 23 | ignore_index=100, |
| 24 | pretraining=False, |
| 25 | detach=True): |
| 26 | super().__init__() |
| 27 | d_model = in_channels |
| 28 | self.ignore_index = ignore_index |
| 29 | self.pretraining = pretraining |
| 30 | self.d_model = d_model |
| 31 | self.detach = detach |
| 32 | self.max_length = max_length + 1 # additional stop token |
| 33 | self.out_channels = out_channels |
| 34 | # -------------------------------------------------------------------------- |
| 35 | # decoder specifics |
| 36 | self.proj = nn.Linear(out_channels, d_model, False) |
| 37 | self.token_encoder = PositionalEncoding(dropout=0.1, |
| 38 | dim=d_model, |
| 39 | max_len=self.max_length) |
| 40 | self.pos_encoder = PositionalEncoding(dropout=0.1, |
| 41 | dim=d_model, |
| 42 | max_len=self.max_length) |
| 43 | |
| 44 | self.decoder = nn.ModuleList([ |
| 45 | TransformerBlock( |
| 46 | d_model=d_model, |
| 47 | nhead=nhead, |
| 48 | dim_feedforward=dim_feedforward, |
| 49 | attention_dropout_rate=dropout, |
| 50 | residual_dropout_rate=dropout, |
| 51 | with_self_attn=False, |
| 52 | with_cross_attn=True, |
| 53 | ) for i in range(num_layers) |
| 54 | ]) |
| 55 | |
| 56 | v_mask = torch.empty((1, 1, d_model)) |
| 57 | l_mask = torch.empty((1, 1, d_model)) |
| 58 | self.v_mask = nn.Parameter(v_mask) |
| 59 | self.l_mask = nn.Parameter(l_mask) |
| 60 | torch.nn.init.uniform_(self.v_mask, -0.001, 0.001) |
| 61 | torch.nn.init.uniform_(self.l_mask, -0.001, 0.001) |
| 62 | |
| 63 | v_embeding = torch.empty((1, 1, d_model)) |
| 64 | l_embeding = torch.empty((1, 1, d_model)) |
| 65 | self.v_embeding = nn.Parameter(v_embeding) |
| 66 | self.l_embeding = nn.Parameter(l_embeding) |
| 67 | torch.nn.init.uniform_(self.v_embeding, -0.001, 0.001) |
| 68 | torch.nn.init.uniform_(self.l_embeding, -0.001, 0.001) |
| 69 | self.cls = nn.Linear(d_model, out_channels) |
| 70 | |
| 71 | def forward_decoder(self, q, x, mask=None): |
| 72 | for decoder_layer in self.decoder: |
nothing calls this directly
no test coverage detected