File size: 3,238 Bytes
3d9ba5b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
import torch
import torch.nn as nn
from typing import Tuple
from .attention import BahdanauAttention

class Seq2SeqDecoder(nn.Module):
    """
    Attention-guided GRU Decoder for Seq2Seq Abstractive Summarization.
    """

    def __init__(
        self,
        output_dim: int,
        emb_dim: int,
        enc_hid_dim: int,
        dec_hid_dim: int,
        dropout: float = 0.2,
        pad_idx: int = 0
    ):
        super().__init__()
        self.output_dim = output_dim
        self.emb_dim = emb_dim
        self.enc_hid_dim = enc_hid_dim
        self.dec_hid_dim = dec_hid_dim

        self.attention = BahdanauAttention(enc_hid_dim * 2, dec_hid_dim)
        self.embedding = nn.Embedding(output_dim, emb_dim, padding_idx=pad_idx)
        self.dropout = nn.Dropout(dropout)

        # GRU takes concatenated [embedded token, context vector]
        # context vector size is enc_hid_dim * 2
        self.rnn = nn.GRU((enc_hid_dim * 2) + emb_dim, dec_hid_dim, batch_first=True)

        # Projection layer to vocabulary logits
        self.fc_out = nn.Linear((enc_hid_dim * 2) + dec_hid_dim + emb_dim, output_dim)

    def forward(
        self,
        input_token: torch.Tensor,
        hidden: torch.Tensor,
        encoder_outputs: torch.Tensor,
        mask: torch.Tensor = None
    ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        """
        Args:
            input_token: [batch_size] (current step token index)
            hidden: [batch_size, dec_hid_dim] (previous decoder hidden state)
            encoder_outputs: [batch_size, src_len, enc_hid_dim * 2]
            mask: [batch_size, src_len]
        Returns:
            prediction: [batch_size, output_dim] (vocabulary probability distribution)
            hidden: [batch_size, dec_hid_dim] (updated decoder hidden state)
            a: [batch_size, src_len] (attention weights)
        """
        # input_token: [batch_size, 1]
        input_token = input_token.unsqueeze(1)
        embedded = self.dropout(self.embedding(input_token)) # [batch_size, 1, emb_dim]

        # Calculate attention weights: [batch_size, src_len]
        a = self.attention(hidden, encoder_outputs, mask=mask)

        # Compute context vector: weighted sum of encoder outputs
        # a.unsqueeze(1): [batch_size, 1, src_len]
        # encoder_outputs: [batch_size, src_len, enc_hid_dim * 2]
        # context: [batch_size, 1, enc_hid_dim * 2]
        context = torch.bmm(a.unsqueeze(1), encoder_outputs)

        # Combine embedded input and context vector
        rnn_input = torch.cat((embedded, context), dim=2) # [batch_size, 1, emb_dim + enc_hid_dim * 2]

        # Step RNN: output shape [batch_size, 1, dec_hid_dim]
        output, hidden = self.rnn(rnn_input, hidden.unsqueeze(0))
        hidden = hidden.squeeze(0) # [batch_size, dec_hid_dim]

        # Combine output, context, and embedding for final prediction
        output = output.squeeze(1)   # [batch_size, dec_hid_dim]
        context = context.squeeze(1) # [batch_size, enc_hid_dim * 2]
        embedded = embedded.squeeze(1) # [batch_size, emb_dim]

        prediction = self.fc_out(torch.cat((output, context, embedded), dim=1)) # [batch_size, output_dim]

        return prediction, hidden, a