import math import nltk import numpy as np import pandas as pd import torch from nltk.util import ngrams from transformers import AutoModelForSeq2SeqLM, AutoTokenizer nltk.download("punkt_tab") FINE_TUNED_MODEL_NAME = "mahdimoghaddami/fine_tuned_bart_base" LED_MODEL_NAME = "allenai/led-large-16384-arxiv" DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") _fine_tuned_tokenizer = None _fine_tuned_model = None _led_tokenizer = None _led_model = None def _load_fine_tuned_resources(): global _fine_tuned_model, _fine_tuned_tokenizer if _fine_tuned_model is None or _fine_tuned_tokenizer is None: _fine_tuned_tokenizer = AutoTokenizer.from_pretrained(FINE_TUNED_MODEL_NAME) _fine_tuned_model = AutoModelForSeq2SeqLM.from_pretrained(FINE_TUNED_MODEL_NAME) _fine_tuned_model.to(DEVICE) _fine_tuned_model.eval() return _fine_tuned_tokenizer, _fine_tuned_model def _load_led_resources(): global _led_model, _led_tokenizer if _led_model is None or _led_tokenizer is None: _led_tokenizer = AutoTokenizer.from_pretrained(LED_MODEL_NAME) _led_model = AutoModelForSeq2SeqLM.from_pretrained(LED_MODEL_NAME) _led_model.to(DEVICE) _led_model.eval() return _led_tokenizer, _led_model def preload_models() -> None: """Preload heavy models once at app startup for faster first request.""" _load_fine_tuned_resources() _load_led_resources() def read_dataset(folder_path: str) -> tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]: train_df = pd.read_pickle(folder_path + "train_set.pkl") test_df = pd.read_pickle(folder_path + "test_set.pkl") val_df = pd.read_pickle(folder_path + "val_set.pkl") return train_df, test_df, val_df def baseline_summarize(text, k=5): sentences = nltk.sent_tokenize(text) if k >= len(sentences): return text scores = [len(s) for s in sentences] sorted_idx = np.argsort(scores)[::-1] selected_idx = sorted_idx[:k] summary = " ".join([sentences[i] for i in selected_idx]) return summary def n_gram_summarize(text: str, n: int = 2, k: int = 5) -> str: def train_ngram_model(text, order=n): """ Trains an N-gram language model from a given text. Args: text: The text to train the model on. order: The N-gram order (e.g., bigrams for 2, trigrams for 3). Returns: A dictionary mapping ngrams to their frequencies. """ tokens = nltk.word_tokenize(text) ngram_counts = {} for ngram in ngrams(tokens, order): ngram_counts[ngram] = ngram_counts.get(ngram, 0) + 1 return ngram_counts def calculate_perplexity(sentence, ngram_model, order=n): """ Calculates the perplexity of a sentence using the trained N-gram model. Args: sentence: The sentence to calculate the perplexity for. ngram_model: The dictionary mapping ngrams to their frequencies. Returns: The calculated perplexity score. """ tokens = nltk.word_tokenize(sentence) log_prob = 0 for i in range(len(tokens)): ngram = tuple(tokens[i : i + order]) if ngram not in ngram_model: # Handle unknown ngrams return float("inf") else: log_prob += math.log(ngram_model[ngram]) perplexity = 2 ** (-log_prob / len(tokens)) return perplexity ngram_model = train_ngram_model(text, order=n) sentences = nltk.sent_tokenize(text) perplexities = [ calculate_perplexity(sentence, ngram_model, n) for sentence in sentences ] sorted_idx = np.argsort(perplexities) selected_idx = sorted_idx[:k] summary = " ".join([sentences[i] for i in selected_idx]) return summary def fine_tuned_summarize(text: str) -> str: tokenizer, model = _load_fine_tuned_resources() # Keep enough context without hitting extreme memory usage. prompted_txt = f"summarize the following: {text}" inputs = tokenizer( prompted_txt, return_tensors="pt", max_length=768, truncation=True ) inputs = {k: v.to(DEVICE) for k, v in inputs.items()} with torch.inference_mode(): summary_ids = model.generate( input_ids=inputs["input_ids"], attention_mask=inputs["attention_mask"], max_new_tokens=300, min_new_tokens=40, num_beams=5, no_repeat_ngram_size=3, length_penalty=1.0, early_stopping=True, ) return tokenizer.decode(summary_ids[0], skip_special_tokens=True) def llm_summarize(text: str) -> str: tokenizer, model = _load_led_resources() prompted_txt = f"summarize the following: {text}" inputs = tokenizer( prompted_txt, return_tensors="pt", max_length=8192, truncation=True ) # LED requires global attention on the first token () global_attention_mask = torch.zeros_like(inputs["attention_mask"]) global_attention_mask[:, 0] = 1 # Move inputs to device input_ids = inputs["input_ids"].to(DEVICE) attention_mask = inputs["attention_mask"].to(DEVICE) global_attention_mask = global_attention_mask.to(DEVICE) # Generate summary with torch.inference_mode(): summary_ids = model.generate( input_ids, attention_mask=attention_mask, global_attention_mask=global_attention_mask, min_new_tokens=50, max_new_tokens=300, num_beams=4, length_penalty=1.0, early_stopping=True, ) # Decode and return return tokenizer.decode(summary_ids[0], skip_special_tokens=True) def summarize_text(model_name: str, num_sentences: int, text: str) -> str: if model_name == "Baseline": return baseline_summarize(text, k=num_sentences) if model_name == "N-gram": return n_gram_summarize(text, k=num_sentences) if model_name == "BART-base": return fine_tuned_summarize(text) return llm_summarize(text)