Views
No views yet
1import librosa
2import torch
3import torchaudio
4from datasets import load_dataset
5from transformers import Wav2Vec2ForCTC, Wav2Vec2Processor
6
7test_dataset = load_dataset("common_voice", "ar", split="test[:2%]")
8
9processor = Wav2Vec2Processor.from_pretrained("kmfoda/wav2vec2-large-xlsr-arabic")
10model = Wav2Vec2ForCTC.from_pretrained("kmfoda/wav2vec2-large-xlsr-arabic")
11
12resamplers = { # all three sampling rates exist in test split
13 48000: torchaudio.transforms.Resample(48000, 16000),
14 44100: torchaudio.transforms.Resample(44100, 16000),
15 32000: torchaudio.transforms.Resample(32000, 16000),
16}
17
18def prepare_example(example):
19 speech, sampling_rate = torchaudio.load(example["path"])
20 example["speech"] = resamplers[sampling_rate](speech).squeeze().numpy()
21 return example
22
23test_dataset = test_dataset.map(prepare_example)
24
25inputs = processor(test_dataset["speech"][:2], sampling_rate=16_000, return_tensors="pt", padding=True)
26
27with torch.no_grad():
28 logits = model(inputs.input_values, attention_mask=inputs.attention_mask).logits
29
30predicted_ids = torch.argmax(logits, dim=-1)
31
32print("Prediction:", processor.batch_decode(predicted_ids))
33print("Reference:", test_dataset["sentence"][:2])1import librosa
2import torch
3import torchaudio
4from datasets import load_dataset, load_metric
5from transformers import Wav2Vec2ForCTC, Wav2Vec2Processor
6import re
7
8test_dataset = load_dataset("common_voice", "ar", split="test")
9wer = load_metric("wer")
10processor = Wav2Vec2Processor.from_pretrained("kmfoda/wav2vec2-large-xlsr-arabic")
11model = Wav2Vec2ForCTC.from_pretrained("kmfoda/wav2vec2-large-xlsr-arabic")
12model.to("cuda")
13
14chars_to_ignore_regex = '[\,\?\.\!\-\;\:\"\“\؟\_\؛\ـ\—]'
15
16resamplers = { # all three sampling rates exist in test split
17 48000: torchaudio.transforms.Resample(48000, 16000),
18 44100: torchaudio.transforms.Resample(44100, 16000),
19 32000: torchaudio.transforms.Resample(32000, 16000),
20}
21
22def prepare_example(example):
23 speech, sampling_rate = torchaudio.load(example["path"])
24 example["speech"] = resamplers[sampling_rate](speech).squeeze().numpy()
25 return example
26
27test_dataset = test_dataset.map(prepare_example)
28
29# Preprocessing the datasets.
30# We need to read the audio files as arrays
31def evaluate(batch):
32 inputs = processor(batch["speech"], sampling_rate=16_000, return_tensors="pt", padding=True)
33
34 with torch.no_grad():
35 logits = model(inputs.input_values.to("cuda"), attention_mask=inputs.attention_mask.to("cuda")).logits
36
37 pred_ids = torch.argmax(logits, dim=-1)
38 batch["pred_strings"] = processor.batch_decode(pred_ids)
39 return batch
40
41result = test_dataset.map(evaluate, batched=True, batch_size=8)
42
43print("WER: {:2f}".format(100 * wer.compute(predictions=result["pred_strings"], references=result["sentence"])))
44train, validation datasets were used for training.