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