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-faroese-100h"
7processor = Wav2Vec2Processor.from_pretrained(MODEL_NAME)
8model = Wav2Vec2ForCTC.from_pretrained(MODEL_NAME)
9
10#Load the dataset
11from datasets import load_dataset, Audio
12ds=load_dataset("carlosdanielhernandezmena/ravnursson_asr",split='test')
13
14#Downsample to 16kHz
15ds = ds.cast_column("audio", Audio(sampling_rate=16_000))
16
17#Process the dataset
18def prepare_dataset(batch):
19 audio = batch["audio"]
20 #Batched output is "un-batched" to ensure mapping is correct
21 batch["input_values"] = processor(audio["array"], sampling_rate=audio["sampling_rate"]).input_values[0]
22 with processor.as_target_processor():
23 batch["labels"] = processor(batch["normalized_text"]).input_ids
24 return batch
25ds = ds.map(prepare_dataset, remove_columns=ds.column_names,num_proc=1)
26
27#Define the evaluation metric
28import numpy as np
29from evaluate import load
30wer_metric = load("wer")
31def compute_metrics(pred):
32 pred_logits = pred.predictions
33 pred_ids = np.argmax(pred_logits, axis=-1)
34 pred.label_ids[pred.label_ids == -100] = processor.tokenizer.pad_token_id
35 pred_str = processor.batch_decode(pred_ids)
36 #We do not want to group tokens when computing the metrics
37 label_str = processor.batch_decode(pred.label_ids, group_tokens=False)
38 wer = wer_metric.compute(predictions=pred_str, references=label_str)
39 return {"wer": wer}
40
41#Do the evaluation (with batch_size=1)
42model = model.to(torch.device("cuda"))
43def map_to_result(batch):
44 with torch.no_grad():
45 input_values = torch.tensor(batch["input_values"], device="cuda").unsqueeze(0)
46 logits = model(input_values).logits
47 pred_ids = torch.argmax(logits, dim=-1)
48 batch["pred_str"] = processor.batch_decode(pred_ids)[0]
49 batch["sentence"] = processor.decode(batch["labels"], group_tokens=False)
50 return batch
51results = ds.map(map_to_result,remove_columns=ds.column_names)
52
53#Compute the overall WER now.
54print("Test WER: {:.3f}".format(wer_metric.compute(predictions=results["pred_str"], references=results["sentence"])))
551@misc{mena2022xlrs53faroese,
2 title={Acoustic Model in Faroese: wav2vec2-large-xlsr-53-faroese-100h.},
3 author={Hernandez Mena, Carlos Daniel},
4 url={https://huggingface.co/carlosdanielhernandezmena/wav2vec2-large-xlsr-53-faroese-100h},
5 year={2022}
6}