import os import random import torch import torch.nn as nn from typing import List, Dict, Tuple, Optional from nlp_core.vocabulary import Vocabulary from .encoder import Seq2SeqEncoder from .decoder import Seq2SeqDecoder class Seq2SeqSummarizer(nn.Module): """ End-to-End Sequence-to-Sequence Abstractive Summarization Model with Bahdanau Attention, Greedy Decoding, and Beam Search. """ def __init__( self, src_vocab: Vocabulary, trg_vocab: Vocabulary, emb_dim: int = 256, enc_hid_dim: int = 512, dec_hid_dim: int = 512, dropout: float = 0.3, device: str = "cpu" ): super().__init__() self.src_vocab = src_vocab self.trg_vocab = trg_vocab self.device = torch.device(device if torch.cuda.is_available() and device == "cuda" else "cpu") self.encoder = Seq2SeqEncoder( input_dim=len(src_vocab), emb_dim=emb_dim, enc_hid_dim=enc_hid_dim, dec_hid_dim=dec_hid_dim, dropout=dropout, pad_idx=src_vocab.PAD_IDX ) self.decoder = Seq2SeqDecoder( output_dim=len(trg_vocab), emb_dim=emb_dim, enc_hid_dim=enc_hid_dim, dec_hid_dim=dec_hid_dim, dropout=dropout, pad_idx=trg_vocab.PAD_IDX ) self.to(self.device) def forward( self, src: torch.Tensor, trg: torch.Tensor, teacher_forcing_ratio: float = 0.5 ) -> torch.Tensor: """ Forward pass during model training. Args: src: [batch_size, src_len] trg: [batch_size, trg_len] teacher_forcing_ratio: Probability of using ground-truth token as next input Returns: outputs: [batch_size, trg_len, trg_vocab_size] """ batch_size = src.shape[0] trg_len = trg.shape[1] trg_vocab_size = self.decoder.output_dim # Tensor to store decoder outputs outputs = torch.zeros(batch_size, trg_len, trg_vocab_size, device=self.device) # Source mask (1 for valid token, 0 for pad) mask = (src != self.src_vocab.PAD_IDX).to(self.device) # Encode source sequence encoder_outputs, hidden = self.encoder(src) # First input to the decoder is the token input_token = trg[:, 0] for t in range(1, trg_len): output, hidden, _ = self.decoder(input_token, hidden, encoder_outputs, mask=mask) outputs[:, t, :] = output # Decide whether to use teacher forcing teacher_force = random.random() < teacher_forcing_ratio top1 = output.argmax(1) input_token = trg[:, t] if teacher_force else top1 return outputs def summarize_greedy( self, src_tokens: List[str], max_len: int = 50 ) -> Tuple[List[str], torch.Tensor]: """ Generates summary tokens using Greedy Search decoding. """ self.eval() with torch.no_grad(): src_indices = self.src_vocab.encode(src_tokens, add_sos=True, add_eos=True) src_tensor = torch.tensor([src_indices], dtype=torch.long, device=self.device) mask = (src_tensor != self.src_vocab.PAD_IDX).to(self.device) encoder_outputs, hidden = self.encoder(src_tensor) trg_indices = [self.trg_vocab.SOS_IDX] attentions = [] for _ in range(max_len): input_token = torch.tensor([trg_indices[-1]], dtype=torch.long, device=self.device) output, hidden, attention = self.decoder(input_token, hidden, encoder_outputs, mask=mask) attentions.append(attention.squeeze(0)) predicted_idx = output.argmax(1).item() if predicted_idx == self.trg_vocab.EOS_IDX: break trg_indices.append(predicted_idx) summary_tokens = self.trg_vocab.decode(trg_indices[1:], remove_special_tokens=True) attn_tensor = torch.stack(attentions) if attentions else torch.empty(0) return summary_tokens, attn_tensor def summarize_beam( self, src_tokens: List[str], beam_width: int = 3, max_len: int = 50, length_penalty: float = 0.7 ) -> List[str]: """ Generates summary tokens using Beam Search decoding with length penalty. """ self.eval() with torch.no_grad(): src_indices = self.src_vocab.encode(src_tokens, add_sos=True, add_eos=True) src_tensor = torch.tensor([src_indices], dtype=torch.long, device=self.device) mask = (src_tensor != self.src_vocab.PAD_IDX).to(self.device) encoder_outputs, hidden = self.encoder(src_tensor) # Beam state: (log_prob, [token_indices], hidden_state) beams = [(0.0, [self.trg_vocab.SOS_IDX], hidden)] completed_beams = [] for _ in range(max_len): new_beams = [] for score, seq, h in beams: if seq[-1] == self.trg_vocab.EOS_IDX: completed_beams.append((score, seq)) continue input_token = torch.tensor([seq[-1]], dtype=torch.long, device=self.device) output, new_h, _ = self.decoder(input_token, h, encoder_outputs, mask=mask) # Log softmax over output log_probs = torch.log_softmax(output, dim=1).squeeze(0) topk_log_probs, topk_indices = torch.topk(log_probs, beam_width) for k in range(beam_width): next_token = topk_indices[k].item() next_score = score + topk_log_probs[k].item() new_beams.append((next_score, seq + [next_token], new_h)) if not new_beams: break # Sort and keep top beam_width # Apply length penalty to ranking score: score / (length ** length_penalty) def rank_key(b): seq_len = len(b[1]) lp = (seq_len ** length_penalty) if seq_len > 0 else 1.0 return b[0] / lp new_beams.sort(key=rank_key, reverse=True) beams = new_beams[:beam_width] if len(completed_beams) >= beam_width: break if not completed_beams: completed_beams = [(score, seq) for score, seq, _ in beams] # Choose best completed beam best_seq = max(completed_beams, key=lambda x: x[0] / (len(x[1]) ** length_penalty))[1] summary_tokens = self.trg_vocab.decode(best_seq[1:], remove_special_tokens=True) return summary_tokens def save_checkpoint(self, filepath: str): """Saves model weights and vocabulary configuration.""" os.makedirs(os.path.dirname(filepath), exist_ok=True) checkpoint = { "state_dict": self.state_dict(), "src_vocab": self.src_vocab.word2idx, "trg_vocab": self.trg_vocab.word2idx, "config": { "emb_dim": self.encoder.emb_dim, "enc_hid_dim": self.encoder.enc_hid_dim, "dec_hid_dim": self.decoder.dec_hid_dim, } } torch.save(checkpoint, filepath) @classmethod def load_checkpoint(cls, filepath: str, device: str = "cpu") -> 'Seq2SeqSummarizer': """Loads model weights and configuration from a saved checkpoint.""" checkpoint = torch.load(filepath, map_location=device) src_vocab = Vocabulary() src_vocab.word2idx = checkpoint["src_vocab"] src_vocab.idx2word = {int(idx): w for w, idx in src_vocab.word2idx.items()} trg_vocab = Vocabulary() trg_vocab.word2idx = checkpoint["trg_vocab"] trg_vocab.idx2word = {int(idx): w for w, idx in trg_vocab.word2idx.items()} config = checkpoint.get("config", {}) model = cls( src_vocab=src_vocab, trg_vocab=trg_vocab, emb_dim=config.get("emb_dim", 256), enc_hid_dim=config.get("enc_hid_dim", 512), dec_hid_dim=config.get("dec_hid_dim", 512), device=device ) model.load_state_dict(checkpoint["state_dict"]) return model