Views
No views yet
1#!/usr/bin/env python3
2from transformers import Wav2Vec2Processor, Wav2Vec2ForCTC
3from datasets import load_dataset
4import torchaudio
5import torch
6
7# resample audio
8
9# load model & processor
10model = Wav2Vec2ForCTC.from_pretrained("facebook/wav2vec2-base-10k-voxpopuli-ft-hu")
11processor = Wav2Vec2Processor.from_pretrained("facebook/wav2vec2-base-10k-voxpopuli-ft-hu")
12
13# load dataset
14ds = load_dataset("common_voice", "hu", split="validation[:1%]")
15
16# common voice does not match target sampling rate
17common_voice_sample_rate = 48000
18target_sample_rate = 16000
19
20resampler = torchaudio.transforms.Resample(common_voice_sample_rate, target_sample_rate)
21
22
23# define mapping fn to read in sound file and resample
24def map_to_array(batch):
25 speech, _ = torchaudio.load(batch["path"])
26 speech = resampler(speech)
27 batch["speech"] = speech[0]
28 return batch
29
30
31# load all audio files
32ds = ds.map(map_to_array)
33
34# run inference on the first 5 data samples
35inputs = processor(ds[:5]["speech"], sampling_rate=target_sample_rate, return_tensors="pt", padding=True)
36
37# inference
38logits = model(**inputs).logits
39predicted_ids = torch.argmax(logits, axis=-1)
40
41print(processor.batch_decode(predicted_ids))