Spaces:
Running on Zero
Running on Zero
Download models/abstractive/seq2seq_model.py from fady21/bilingual-summarizer-api: direct link, hf CLI and curl.
- Browser
- Download file 8.52 kB
-
https://huggingface.co/spaces/fady21/bilingual-summarizer-api/resolve/main/models/abstractive/seq2seq_model.py
- Command line
-
hf download hf://spaces/fady21/bilingual-summarizer-api/models/abstractive/seq2seq_model.py
-
curl -L -o seq2seq_model.py https://huggingface.co/spaces/fady21/bilingual-summarizer-api/resolve/main/models/abstractive/seq2seq_model.py
8.52 kB
| 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 <sos> 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) | |
| 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 | |