Ahmed007/hamsa-tiny-v0.4
09
1import torch2import librosa3from datasets import load_dataset, Audio4from transformers import WhisperProcessor, WhisperFeatureExtractor, WhisperTokenizer, WhisperForConditionalGeneration5from huggingface_hub import login6import argparse7from evaluate import load8 9my_parser = argparse.ArgumentParser()10# my_parser.add_argument("--pal", "-paths_as_labels", action="store_true")11 12my_parser.add_argument("--model_name", "-model_name", type=str, action="store", default = "openai/whisper-tiny")13my_parser.add_argument("--hf_token", "-hf_token", type=str, action="store")14my_parser.add_argument("--dataset_name", "-dataset_name", type=str, action="store", default = "google/fleurs")15my_parser.add_argument("--split", "-split", type=str, action="store", default = "test")16my_parser.add_argument("--subset", "-subset", type=str, action="store")17 18args = my_parser.parse_args()19try:20 login(args.hf_token)21except:22 raise(f"Can't login please set --hf_token {args.hf_token}")23 24 25dataset_name = args.dataset_name 26model_name = args.model_name27subset = args.subset28text_column = "sentence"29if dataset_name == "google/fleurs":30 text_column = "transcription"31 32print(f"Evaluating {args.model_name} on {args.dataset_name} [{subset}]")33 34 35feature_extractor = WhisperFeatureExtractor.from_pretrained(model_name)36model = WhisperForConditionalGeneration.from_pretrained(model_name)37 38test_dataset = load_dataset(dataset_name, subset, split=args.split, use_auth_token=True)39processor = WhisperProcessor.from_pretrained(model_name, language="Arabic", task="transcribe")40tokenizer = WhisperTokenizer.from_pretrained(model_name, language="Arabic", task="transcribe")41test_dataset = test_dataset.cast_column("audio", Audio(sampling_rate=16000))42 43# Preprocessing the datasets.44def prepare_dataset(batch):45 # load and resample audio data from 48 to 16kHz46 audio = batch["audio"]47 48 # compute log-Mel input features from input audio array 49 batch["input_features"] = feature_extractor(audio["array"], sampling_rate=audio["sampling_rate"]).input_features[0]50 51 # encode target text to label ids 52 batch["labels"] = tokenizer(batch[text_column]).input_ids53 return batch54 55test_dataset = test_dataset.map(prepare_dataset)56 57model = model.to("cuda")58model.config.forced_decoder_ids = processor.get_decoder_prompt_ids(language = "ar", task = "transcribe")59 60def map_to_result(batch):61 62 with torch.no_grad():63 input_values = torch.tensor(batch["input_features"], device="cuda").unsqueeze(0)64 pred_ids = model.generate(input_values)65 66 batch["pred_str"] = processor.batch_decode(pred_ids, skip_special_tokens = True)[0]67 batch["text"] = processor.decode(batch["labels"], skip_special_tokens = True)68 69 return batch70results = test_dataset.map(map_to_result)71 72wer = load("wer")73print("Test WER: {:.3f}".format(wer.compute(predictions=results["pred_str"], references=results["text"])))74 