From 23a1580fa914fb8182574f4f057a81a431fda1a2 Mon Sep 17 00:00:00 2001 From: amitroth Date: Sat, 8 Mar 2025 16:10:54 +0200 Subject: [PATCH] added support for text only metrics, and an implementation of hellaswag --- cli/eval.py | 3 ++ config/metric/hellaswag.yaml | 6 +++ slamkit/metric/textual_metric.py | 74 ++++++++++++++++++++++++++++++++ slamkit/model/speech_lm.py | 13 ++++++ 4 files changed, 96 insertions(+) create mode 100644 config/metric/hellaswag.yaml create mode 100644 slamkit/metric/textual_metric.py diff --git a/cli/eval.py b/cli/eval.py index ec0775a..14c314d 100644 --- a/cli/eval.py +++ b/cli/eval.py @@ -7,6 +7,7 @@ from slamkit.tokeniser import tokeniser_factory from slamkit.metric.generative_metric import generate, asr_perplexity from slamkit.metric.modelling_metric import swuggy, salmon, sblimp, storycloze +from slamkit.metric.textual_metric import hellaswag from slamkit.model import SpeechLM import torch import logging @@ -51,6 +52,8 @@ def main(cfg: DictConfig): res = asr_perplexity(model, path, cfg.batch_size, cfg.metric.whisper_model, cfg.metric.llm_name_or_path, used_token_modality, cfg.metric.prompt_length, cfg.metric.auto_bleu_n, tokeniser.fe_sample_rate, cfg.metric.get("num_files", None), cfg.num_workers, cfg.pin_memory, **cfg.metric.get("generate_kwargs", {})) + elif cfg.metric.metric_type == 'hellaswag': + res = hellaswag(model, path, used_token_modality, mean_nll, cfg.batch_size, cfg.num_workers, cfg.pin_memory, cfg.metric.get("subfolder", False)) else: raise ValueError(f'Unknown metric type: {cfg.metric.metric_type}') if cfg.metric.metric_type != "generate": diff --git a/config/metric/hellaswag.yaml b/config/metric/hellaswag.yaml new file mode 100644 index 0000000..0da518b --- /dev/null +++ b/config/metric/hellaswag.yaml @@ -0,0 +1,6 @@ +defaults: + - default + - _self_ + +metric_type: hellaswag +data_path: //reference/hellaswag diff --git a/slamkit/metric/textual_metric.py b/slamkit/metric/textual_metric.py new file mode 100644 index 0000000..1919868 --- /dev/null +++ b/slamkit/metric/textual_metric.py @@ -0,0 +1,74 @@ +import logging +logger = logging.getLogger(__name__) + +import json +import torch +from pathlib import Path +from torch.utils.data import DataLoader, Dataset +from torch.nn.utils.rnn import pad_sequence +from tqdm import tqdm +import re + + +class HellaSwagDataset(Dataset): + def __init__(self, path): + super().__init__() + self.data = [] + + with open(path, 'r') as file: + self.data = [json.loads(line) for line in file] + + def __len__(self): + return len(self.data) + + def __getitem__(self, idx): + data = self.data[idx] + positive_index = data['label'] + ctx = data["ctx_a"] + " " + data["ctx_b"].capitalize() + query = HellaSwagDataset.preprocess(data["activity_label"] + ": " + ctx) + endings = [HellaSwagDataset.preprocess(ending) for ending in data['endings']] + full_sentences = [query + ending for ending in endings] + + return full_sentences[positive_index:] + full_sentences[:positive_index] + + @staticmethod + def preprocess(text): + text = text.strip() + # NOTE: Brackets are artifacts of the WikiHow dataset portion of HellaSwag. + text = text.replace(" [title]", ". ") + text = re.sub("\\[.*?\\]", "", text) + text = text.replace(" ", " ") + return text + + + +def textual_metric(model, dataset, used_token_modality, mean_nll: bool=True, + batch_size: int = 1, num_workers=8, pin_memory=True): + dl = DataLoader(dataset, batch_size=batch_size ,num_workers=num_workers, pin_memory=pin_memory) + res_list = [] + + counter = 0 + for sample_files in tqdm(dl): + counter +=1 + + with torch.no_grad(): + results = [ + model.text_log_likelihood(sample, used_token_modality=used_token_modality, mean_nll=mean_nll) + for sample in sample_files + ] + + res = (results[0] > torch.stack(results[1:]).max(dim=0).values).int() + res_list.append(res) + + res_list = torch.cat(res_list) + return res_list.float().mean().cpu().item() + + +def hellaswag(model, data_path, used_token_modality, mean_nll=True, + batch_size=1, num_workers=8, pin_memory=True): + dataset = HellaSwagDataset(data_path) + assert len(dataset) > 0, f"no samples found for {data_path}" + res = textual_metric(model, dataset, used_token_modality, mean_nll, batch_size, num_workers, pin_memory) + logging.info(f"HellaSwag: {res:.4f}") + return {'HellaSwag': res} + diff --git a/slamkit/model/speech_lm.py b/slamkit/model/speech_lm.py index 35f5874..7965b13 100644 --- a/slamkit/model/speech_lm.py +++ b/slamkit/model/speech_lm.py @@ -35,6 +35,19 @@ def log_likelihood(self, wavs: torch.Tensor, lens: Optional[torch.Tensor] = None ignore_tokens = self.tokeniser.get_ignore_tokens(used_token_modality) return self.model.log_likelihood(tokens, mean_nll, ignore_tokens) + def text_log_likelihood(self, texts: List[str], mean_nll: bool = True, used_token_modality: Optional[str] = None) -> torch.Tensor: + """ + Given a list of texts, calculate the log likelihood for each sample. + :param texts: A list of strings + :param mean_nll: whether to take mean instead of sum thus cancelling length bias + :param used_token_modality: the tokens modality to use + :return: + """ + tokens = self.tokeniser.string_tokenise(texts, return_tensors='pt', padding=True)['input_ids'].to(self.device) + ignore_tokens = self.tokeniser.get_ignore_tokens(used_token_modality) + return self.model.log_likelihood(tokens, mean_nll, ignore_tokens) + + def generate(self, wavs: torch.Tensor, lens: Optional[torch.Tensor] = None, used_token_modality: Optional[str] = None, remove_prompt=False, **kwargs) -> List[torch.Tensor]: """ Given a batch of wavs zero padded, generate the continuation tokens or audio if a vocoder is present