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-spanish-ep5-944h"
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("ciempiess/ciempiess_test", 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
29wer_metric = load_metric("wer")
30def compute_metrics(pred):
31 pred_logits = pred.predictions
32 pred_ids = np.argmax(pred_logits, axis=-1)
33 pred.label_ids[pred.label_ids == -100] = processor.tokenizer.pad_token_id
34 pred_str = processor.batch_decode(pred_ids)
35 #We do not want to group tokens when computing the metrics
36 label_str = processor.batch_decode(pred.label_ids, group_tokens=False)
37 wer = wer_metric.compute(predictions=pred_str, references=label_str)
38 return {"wer": wer}
39
40#Do the evaluation (with batch_size=1)
41model = model.to(torch.device("cuda"))
42def map_to_result(batch):
43 with torch.no_grad():
44 input_values = torch.tensor(batch["input_values"], device="cuda").unsqueeze(0)
45 logits = model(input_values).logits
46 pred_ids = torch.argmax(logits, dim=-1)
47 batch["pred_str"] = processor.batch_decode(pred_ids)[0]
48 batch["sentence"] = processor.decode(batch["labels"], group_tokens=False)
49 return batch
50results = ds.map(map_to_result,remove_columns=ds.column_names)
51
52#Compute the overall WER now.
53print("Test WER: {:.3f}".format(wer_metric.compute(predictions=results["pred_str"], references=results["sentence"])))1@misc{mena2022xlrs53spanish,
2 title={Acoustic Model in Spanish: wav2vec2-large-xlsr-53-spanish-ep5-944h.},
3 author={Hernandez Mena, Carlos Daniel},
4 url={https://huggingface.co/carlosdanielhernandezmena/wav2vec2-large-xlsr-53-spanish-ep5-944h},
5 year={2022}
6}