Views
No views yet
1import torch
2from transformers import Wav2Vec2Processor
3from transformers import Wav2Vec2ForCTC
4
5#Load the processor and model.
6MODEL_NAME="carlosdanielhernandezmena/wav2vec2-large-xlsr-53-maltese-64h"
7processor = Wav2Vec2Processor.from_pretrained(MODEL_NAME)
8model = Wav2Vec2ForCTC.from_pretrained(MODEL_NAME)
9
10#Load the dataset
11from datasets import load_dataset, load_metric, Audio
12ds=load_dataset("common_voice", "mt", split="test")
13
14#Normalize the transcriptions
15import re
16chars_to_ignore_regex = '[\\,\\?\\.\\!\\\;\\:\\"\\“\\%\\‘\\”\\�\\)\\(\\*)]'
17def remove_special_characters(batch):
18 batch["sentence"] = re.sub(chars_to_ignore_regex, '', batch["sentence"]).lower()
19 return batch
20ds = ds.map(remove_special_characters)
21
22#Downsample to 16kHz
23ds = ds.cast_column("audio", Audio(sampling_rate=16_000))
24
25#Process the dataset
26def prepare_dataset(batch):
27 audio = batch["audio"]
28 #Batched output is "un-batched" to ensure mapping is correct
29 batch["input_values"] = processor(audio["array"], sampling_rate=audio["sampling_rate"]).input_values[0]
30 with processor.as_target_processor():
31 batch["labels"] = processor(batch["sentence"]).input_ids
32 return batch
33ds = ds.map(prepare_dataset, remove_columns=ds.column_names,num_proc=1)
34
35#Define the evaluation metric
36import numpy as np
37wer_metric = load_metric("wer")
38def compute_metrics(pred):
39 pred_logits = pred.predictions
40 pred_ids = np.argmax(pred_logits, axis=-1)
41 pred.label_ids[pred.label_ids == -100] = processor.tokenizer.pad_token_id
42 pred_str = processor.batch_decode(pred_ids)
43 #We do not want to group tokens when computing the metrics
44 label_str = processor.batch_decode(pred.label_ids, group_tokens=False)
45 wer = wer_metric.compute(predictions=pred_str, references=label_str)
46 return {"wer": wer}
47
48#Do the evaluation (with batch_size=1)
49model = model.to(torch.device("cuda"))
50def map_to_result(batch):
51 with torch.no_grad():
52 input_values = torch.tensor(batch["input_values"], device="cuda").unsqueeze(0)
53 logits = model(input_values).logits
54 pred_ids = torch.argmax(logits, dim=-1)
55 batch["pred_str"] = processor.batch_decode(pred_ids)[0]
56 batch["sentence"] = processor.decode(batch["labels"], group_tokens=False)
57 return batch
58results = ds.map(map_to_result,remove_columns=ds.column_names)
59
60#Compute the overall WER now.
61print("Test WER: {:.3f}".format(wer_metric.compute(predictions=results["pred_str"], references=results["sentence"])))
621@misc{mena2022xlrs53maltese,
2 title={Acoustic Model in Maltese: wav2vec2-large-xlsr-53-maltese-64h.},
3 author={Hernandez Mena, Carlos Daniel},
4 url={https://huggingface.co/carlosdanielhernandezmena/wav2vec2-large-xlsr-53-maltese-64h},
5 year={2022}
6}