-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathEncoderDecoder.py
More file actions
34 lines (29 loc) · 1.34 KB
/
Copy pathEncoderDecoder.py
File metadata and controls
34 lines (29 loc) · 1.34 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
import torch.nn as nn
class EncoderDecoder(nn.Module):
"""
A standard Encoder-Decoder architecture. Base for this and many
other models.
"""
def __init__(self, encoder, decoder, src_embed, tgt_embed, generator):
super(EncoderDecoder, self).__init__()
self.encoder = encoder
self.decoder = decoder
self.src_embed = src_embed
self.tgt_embed = tgt_embed
self.generator = generator
def forward(self, src, tgt, src_mask, tgt_mask):
"Take in and process masked src and target sequences."
return self.decode(self.encode(src, src_mask), src_mask,
tgt, tgt_mask)
def encode(self, src, src_mask):
#print("self.src_mask.shape",src_mask.shape)
#print("self.src_mask",src_mask)
#print("self.src_embed(src).shape",self.src_embed(src).shape)
#print("self.src_embed(src)",self.src_embed(src))
return self.encoder(self.src_embed(src), src_mask)
def decode(self, memory, src_mask, tgt, tgt_mask):
#print("self.tgt_mask.shape",tgt_mask.shape)
#print("self.tgt_mask",tgt_mask)
#print("self.tgt_embed(tgt).shape",self.tgt_embed(tgt).shape)
#print("self.tgt_embed(tgt)",self.tgt_embed(tgt))
return self.decoder(self.tgt_embed(tgt), memory, src_mask, tgt_mask)