import torch import torch.nn as nn from typing import Tuple class Seq2SeqEncoder(nn.Module): """ Bidirectional GRU Encoder for Seq2Seq Abstractive Summarization. """ def __init__( self, input_dim: int, emb_dim: int, enc_hid_dim: int, dec_hid_dim: int, n_layers: int = 1, dropout: float = 0.2, pad_idx: int = 0 ): super().__init__() self.input_dim = input_dim self.emb_dim = emb_dim self.enc_hid_dim = enc_hid_dim self.dec_hid_dim = dec_hid_dim self.n_layers = n_layers self.embedding = nn.Embedding(input_dim, emb_dim, padding_idx=pad_idx) self.dropout = nn.Dropout(dropout) self.rnn = nn.GRU( emb_dim, enc_hid_dim, num_layers=n_layers, bidirectional=True, batch_first=True, dropout=dropout if n_layers > 1 else 0.0 ) # Linear projection to transform bidirectional hidden states to decoder hidden size self.fc = nn.Linear(enc_hid_dim * 2, dec_hid_dim) def forward(self, src: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: """ Args: src: [batch_size, src_len] Returns: outputs: [batch_size, src_len, enc_hid_dim * 2] (Contextual token representations) hidden: [batch_size, dec_hid_dim] (Initial hidden state for decoder) """ # embedded: [batch_size, src_len, emb_dim] embedded = self.dropout(self.embedding(src)) # outputs: [batch_size, src_len, enc_hid_dim * 2] # hidden: [n_layers * 2, batch_size, enc_hid_dim] outputs, hidden = self.rnn(embedded) # Concatenate final forward and backward hidden states from the top layer # hidden[-2, :, :] is forward, hidden[-1, :, :] is backward final_hidden = torch.cat((hidden[-2, :, :], hidden[-1, :, :]), dim=1) # Project to decoder hidden dimension: [batch_size, dec_hid_dim] dec_init_hidden = torch.tanh(self.fc(final_hidden)) return outputs, dec_init_hidden