| 17 | return torch.matmul(p_attn, value), p_attn |
| 18 | |
| 19 | class EncoderDecoder(pl.LightningModule): |
| 20 | def __init__(self, encoder, decoder, src_embed, tgt_embed, generator): |
| 21 | super().__init__() |
| 22 | self.encoder = encoder |
| 23 | self.decoder = decoder |
| 24 | self.src_embed = src_embed |
| 25 | self.tgt_embed = tgt_embed |
| 26 | self.generator = generator |
| 27 | |
| 28 | def encode(self, src, src_mask): |
| 29 | return self.encoder(self.src_embed(src), src_mask) |
| 30 | |
| 31 | def decode(self, memory, src_mask, tgt, tgt_mask): |
| 32 | return self.generator(self.decoder(self.tgt_embed(tgt), memory, src_mask, tgt_mask)) |
| 33 | |
| 34 | def forward(self, src, tgt, src_mask, tgt_mask): |
| 35 | return self.decode(self.encode(src, src_mask), src_mask, tgt, tgt_mask) |
| 36 | |
| 37 | class Encoder(pl.LightningModule): |
| 38 | def __init__(self, layer, N): |