fin09
Deploy Bilingual Summarization NLP Suite with Git LFS
3d9ba5b
Raw History Blame Contribute Delete
3.24 kB
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