Views
No views yet
train splits of Common Voice
and Arabic Speech Corpus.
When using this model, make sure that your speech input is sampled at 16kHz.1import torch
2import torchaudio
3from datasets import load_dataset
4from lang_trans.arabic import buckwalter
5from transformers import Wav2Vec2ForCTC, Wav2Vec2Processor
6
7dataset = load_dataset("common_voice", "ar", split="test[:10]")
8resamplers = { # all three sampling rates exist in test split
9 48000: torchaudio.transforms.Resample(48000, 16000),
10 44100: torchaudio.transforms.Resample(44100, 16000),
11 32000: torchaudio.transforms.Resample(32000, 16000),
12}
13
14def prepare_example(example):
15 speech, sampling_rate = torchaudio.load(example["path"])
16 example["speech"] = resamplers[sampling_rate](speech).squeeze().numpy()
17 return example
18
19dataset = dataset.map(prepare_example)
20processor = Wav2Vec2Processor.from_pretrained("elgeish/wav2vec2-large-xlsr-53-arabic")
21model = Wav2Vec2ForCTC.from_pretrained("elgeish/wav2vec2-large-xlsr-53-arabic").eval()
22
23def predict(batch):
24 inputs = processor(batch["speech"], sampling_rate=16000, return_tensors="pt", padding=True)
25 with torch.no_grad():
26 predicted = torch.argmax(model(inputs.input_values).logits, dim=-1)
27 predicted[predicted == -100] = processor.tokenizer.pad_token_id # see fine-tuning script
28 batch["predicted"] = processor.tokenizer.batch_decode(predicted)
29 return batch
30
31dataset = dataset.map(predict, batched=True, batch_size=1, remove_columns=["speech"])
32
33for reference, predicted in zip(dataset["sentence"], dataset["predicted"]):
34 print("reference:", reference)
35 print("predicted:", buckwalter.untrans(predicted))
36 print("--")reference: ألديك قلم ؟
predicted: هلديك قالر
--
reference: ليست هناك مسافة على هذه الأرض أبعد من يوم أمس.
predicted: ليست نالك مسافة على هذه الأرض أبعد من يوم أمس
--
reference: إنك تكبر المشكلة.
predicted: إنك تكبر المشكلة
--
reference: يرغب أن يلتقي بك.
predicted: يرغب أن يلتقي بك
--
reference: إنهم لا يعرفون لماذا حتى.
predicted: إنهم لا يعرفون لماذا حتى
--
reference: سيسعدني مساعدتك أي وقت تحب.
predicted: سيسئدني مساعد سكرأي وقت تحب
--
reference: أَحَبُّ نظريّة علمية إليّ هي أن حلقات زحل مكونة بالكامل من الأمتعة المفقودة.
predicted: أحب ناضريةً علمية إلي هي أنحل قتزح المكونا بالكامل من الأمت عن المفقودة
--
reference: سأشتري له قلماً.
predicted: سأشتري له قلما
--
reference: أين المشكلة ؟
predicted: أين المشكل
--
reference: وَلِلَّهِ يَسْجُدُ مَا فِي السَّمَاوَاتِ وَمَا فِي الْأَرْضِ مِنْ دَابَّةٍ وَالْمَلَائِكَةُ وَهُمْ لَا يَسْتَكْبِرُونَ
predicted: ولله يسجد ما في السماوات وما في الأرض من دابة والملائكة وهم لا يستكبرون
--1import jiwer
2import torch
3import torchaudio
4from datasets import load_dataset
5from lang_trans.arabic import buckwalter
6from transformers import set_seed, Wav2Vec2ForCTC, Wav2Vec2Processor
7
8set_seed(42)
9test_split = load_dataset("common_voice", "ar", split="test")
10resamplers = { # all three sampling rates exist in test split
11 48000: torchaudio.transforms.Resample(48000, 16000),
12 44100: torchaudio.transforms.Resample(44100, 16000),
13 32000: torchaudio.transforms.Resample(32000, 16000),
14}
15
16def prepare_example(example):
17 speech, sampling_rate = torchaudio.load(example["path"])
18 example["speech"] = resamplers[sampling_rate](speech).squeeze().numpy()
19 return example
20
21test_split = test_split.map(prepare_example)
22processor = Wav2Vec2Processor.from_pretrained("elgeish/wav2vec2-large-xlsr-53-arabic")
23model = Wav2Vec2ForCTC.from_pretrained("elgeish/wav2vec2-large-xlsr-53-arabic").to("cuda").eval()
24
25def predict(batch):
26 inputs = processor(batch["speech"], sampling_rate=16000, return_tensors="pt", padding=True)
27 with torch.no_grad():
28 predicted = torch.argmax(model(inputs.input_values.to("cuda")).logits, dim=-1)
29 predicted[predicted == -100] = processor.tokenizer.pad_token_id # see fine-tuning script
30 batch["predicted"] = processor.batch_decode(predicted)
31 return batch
32
33test_split = test_split.map(predict, batched=True, batch_size=16, remove_columns=["speech"])
34transformation = jiwer.Compose([
35 # normalize some diacritics, remove punctuation, and replace Persian letters with Arabic ones
36 jiwer.SubstituteRegexes({
37 r'[auiFNKo\~_،؟»\?;:\-,\.؛«!"]': "", "\u06D6": "",
38 r"[\|\{]": "A", "p": "h", "ک": "k", "ی": "y"}),
39 # default transformation below
40 jiwer.RemoveMultipleSpaces(),
41 jiwer.Strip(),
42 jiwer.SentencesToListOfWords(),
43 jiwer.RemoveEmptyStrings(),
44])
45metrics = jiwer.compute_measures(
46 truth=[buckwalter.trans(s) for s in test_split["sentence"]], # Buckwalter transliteration
47 hypothesis=test_split["predicted"],
48 truth_transform=transformation,
49 hypothesis_transform=transformation,
50)
51print(f"WER: {metrics['wer']:.2%}")">" maps to "أ").
The lang-trans package is used to convert (transliterate) Arabic abjad.train split of the Arabic Speech Corpus dataset;
the test split was used for model selection; the resulting model at this point is saved as elgeish/wav2vec2-large-xlsr-53-levantine-arabic.train split of the Common Voice dataset;
the validation split was used for model selection;
training was stopped to meet the deadline of Fine-Tune-XLSR Week:
this model is the checkpoint at 100k steps and a validation WER of 23.39%.attention_mask in model input, which is recommended here.
Also, exploring data augmentation using datasets used to train models listed here.