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