AI_Summarizer / src /textSummarizer /components /model_evaluation.py
Jeevant10's picture
completed project AI Summarizer
65db57d
Raw
History Blame Contribute Delete
3.02 kB
import os
import evaluate
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
from textSummarizer.entity import ModelEvaluationConfig
from datasets import load_from_disk
import torch
import pandas as pd
from tqdm import tqdm
import logging
logger = logging.getLogger(__name__)
class ModelEvaluation:
def __init__(self, config: ModelEvaluationConfig):
self.config = config
def generate_batch_sized_chunks(self, list_of_elements, batch_size):
for i in range(0, len(list_of_elements), batch_size):
yield list_of_elements[i : i + batch_size]
def calculate_metric_on_test_ds(self, dataset, metric, model, tokenizer,
batch_size=16, device="cuda" if torch.cuda.is_available() else "cpu",
column_text="article", column_summary="highlights"):
article_batches = list(self.generate_batch_sized_chunks(dataset[column_text], batch_size))
target_batches = list(self.generate_batch_sized_chunks(dataset[column_summary], batch_size))
for article_batch, target_batch in tqdm(
zip(article_batches, target_batches), total=len(article_batches)):
inputs = tokenizer(article_batch, max_length=1024, truncation=True,
padding="max_length", return_tensors="pt")
summaries = model.generate(input_ids=inputs["input_ids"].to(device),
attention_mask=inputs["attention_mask"].to(device),
length_penalty=0.8, num_beams=8, max_length=128)
decoded_summaries = [tokenizer.decode(s, skip_special_tokens=True, clean_up_tokenization_spaces=True)
for s in summaries]
metric.add_batch(predictions=decoded_summaries, references=target_batch)
score = metric.compute()
return score
def evaluate(self):
logger.info("Loading tokenizer and model...")
device = "cuda" if torch.cuda.is_available() else "cpu"
tokenizer = AutoTokenizer.from_pretrained(self.config.tokenizer_path)
model_pegasus = AutoModelForSeq2SeqLM.from_pretrained(self.config.model_path).to(device)
logger.info("Loading dataset...")
dataset_samsum_pt = load_from_disk(self.config.data_path)
rouge_names = ["rouge1", "rouge2", "rougeL", "rougeLsum"]
rouge_metric = evaluate.load('rouge')
logger.info("Starting evaluation...")
score = self.calculate_metric_on_test_ds(
dataset_samsum_pt['test'][0:10], rouge_metric, model_pegasus, tokenizer, batch_size=2,
column_text='dialogue', column_summary='summary'
)
rouge_dict = {rn: score[rn] for rn in rouge_names}
df = pd.DataFrame(rouge_dict, index=['pegasus'])
logger.info(f"Saving metrics to {self.config.metric_file_name}")
df.to_csv(self.config.metric_file_name, index=False)