MCPcopy Create free account
hub / github.com/TEA-Lab/TwoByTwo / EncoderDecoder

Class EncoderDecoder

src/shape_assembly/models/train/transformer.py:19–35  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

17 return torch.matmul(p_attn, value), p_attn
18
19class 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
37class Encoder(pl.LightningModule):
38 def __init__(self, layer, N):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected