Views
No views yet
1import torchaudio
2from datasets import load_dataset, load_metric
3from transformers import (
4 Wav2Vec2ForCTC,
5 Wav2Vec2Processor,
6)
7import torch
8import re
9import sys
10
11model_name = "facebook/wav2vec2-large-xlsr-53-spanish"
12device = "cuda"
13
14chars_to_ignore_regex = '[\,\?\.\!\-\;\:\"]' # noqa: W605
15
16model = Wav2Vec2ForCTC.from_pretrained(model_name).to(device)
17processor = Wav2Vec2Processor.from_pretrained(model_name)
18
19ds = load_dataset("common_voice", "es", split="test", data_dir="./cv-corpus-6.1-2020-12-11")
20
21resampler = torchaudio.transforms.Resample(orig_freq=48_000, new_freq=16_000)
22
23def map_to_array(batch):
24 speech, _ = torchaudio.load(batch["path"])
25 batch["speech"] = resampler.forward(speech.squeeze(0)).numpy()
26 batch["sampling_rate"] = resampler.new_freq
27 batch["sentence"] = re.sub(chars_to_ignore_regex, '', batch["sentence"]).lower().replace("’", "'")
28 return batch
29
30ds = ds.map(map_to_array)
31
32def map_to_pred(batch):
33 features = processor(batch["speech"], sampling_rate=batch["sampling_rate"][0], padding=True, return_tensors="pt")
34 input_values = features.input_values.to(device)
35 attention_mask = features.attention_mask.to(device)
36 with torch.no_grad():
37 logits = model(input_values, attention_mask=attention_mask).logits
38 pred_ids = torch.argmax(logits, dim=-1)
39 batch["predicted"] = processor.batch_decode(pred_ids)
40 batch["target"] = batch["sentence"]
41 return batch
42
43result = ds.map(map_to_pred, batched=True, batch_size=16, remove_columns=list(ds.features.keys()))
44
45wer = load_metric("wer")
46
47print(wer.compute(predictions=result["predicted"], references=result["target"]))