Views
No views yet
1import torch
2import torchaudio
3from datasets import load_dataset, load_metric
4from transformers import (
5 Wav2Vec2ForCTC,
6 Wav2Vec2Processor,
7)
8import re
9
10model_name = "Ilyes/wav2vec2-large-xlsr-53-french"
11
12device = "cpu" # "cuda"
13
14model = Wav2Vec2ForCTC.from_pretrained(model_name).to(device)
15processor = Wav2Vec2Processor.from_pretrained(model_name)
16
17ds = load_dataset("common_voice", "fr", split="test", cache_dir="./data/fr")
18
19chars_to_ignore_regex = '[\,\?\.\!\;\:\"\“\%\‘\”\�\‘\’\’\’\‘\…\·\!\ǃ\?\«\‹\»\›“\”\\ʿ\ʾ\„\∞\\|\.\,\;\:\*\—\–\─\―\_\/\:\ː\;\,\=\«\»\→]'
20def map_to_array(batch):
21 speech, _ = torchaudio.load(batch["path"])
22 batch["speech"] = resampler.forward(speech.squeeze(0)).numpy()
23 batch["sampling_rate"] = resampler.new_freq
24 batch["sentence"] = re.sub(chars_to_ignore_regex, '', batch["sentence"]).lower().replace("’", "'")
25 return batch
26resampler = torchaudio.transforms.Resample(48_000, 16_000)
27
28ds = ds.map(map_to_array)
29
30def map_to_pred(batch):
31 features = processor(batch["speech"], sampling_rate=batch["sampling_rate"][0], padding=True, return_tensors="pt")
32 input_values = features.input_values.to(device)
33 attention_mask = features.attention_mask.to(device)
34 with torch.no_grad():
35 logits = model(input_values, attention_mask=attention_mask).logits
36 pred_ids = torch.argmax(logits, dim=-1)
37 batch["predicted"] = processor.batch_decode(pred_ids)
38 batch["target"] = batch["sentence"]
39 return batch
40
41result = ds.map(map_to_pred, batched=True, batch_size=16, remove_columns=list(ds.features.keys()))
42wer = load_metric("wer")
43print(wer.compute(predictions=result["predicted"], references=result["target"]))