Spaces:
Running on Zero
Running on Zero
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
|