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