Views
No views yet
| Model | Common Voice 9.0 |
|---|---|
| openai/whisper-small | 13.0 |
| openai/whisper-medium | 8.5 |
| openai/whisper-large-v2 | 6.4 |
| Model | Common Voice 11.0 |
|---|---|
| bofenghuang/whisper-small-cv11-german | 11.35 |
| bofenghuang/whisper-medium-cv11-german | 7.05 |
| bofenghuang/whisper-large-v2-cv11-german | 5.76 |
1import torch
2
3from datasets import load_dataset
4from transformers import pipeline
5
6device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
7
8# Load pipeline
9pipe = pipeline("automatic-speech-recognition", model="bofenghuang/whisper-medium-cv11-german", device=device)
10
11# NB: set forced_decoder_ids for generation utils
12pipe.model.config.forced_decoder_ids = pipe.tokenizer.get_decoder_prompt_ids(language="de", task="transcribe")
13
14# Load data
15ds_mcv_test = load_dataset("mozilla-foundation/common_voice_11_0", "de", split="test", streaming=True)
16test_segment = next(iter(ds_mcv_test))
17waveform = test_segment["audio"]
18
19# NB: decoding option
20# limit the maximum number of generated tokens to 225
21pipe.model.config.max_length = 225 + 1
22# sampling
23# pipe.model.config.do_sample = True
24# beam search
25# pipe.model.config.num_beams = 5
26# return
27# pipe.model.config.return_dict_in_generate = True
28# pipe.model.config.output_scores = True
29# pipe.model.config.num_return_sequences = 5
30
31# Run
32generated_sentences = pipe(waveform)["text"]1import torch
2import torchaudio
3
4from datasets import load_dataset
5from transformers import AutoProcessor, AutoModelForSpeechSeq2Seq
6
7device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
8
9# Load model
10model = AutoModelForSpeechSeq2Seq.from_pretrained("bofenghuang/whisper-medium-cv11-german").to(device)
11processor = AutoProcessor.from_pretrained("bofenghuang/whisper-medium-cv11-german", language="german", task="transcribe")
12
13# NB: set forced_decoder_ids for generation utils
14model.config.forced_decoder_ids = processor.get_decoder_prompt_ids(language="de", task="transcribe")
15
16# 16_000
17model_sample_rate = processor.feature_extractor.sampling_rate
18
19# Load data
20ds_mcv_test = load_dataset("mozilla-foundation/common_voice_11_0", "de", split="test", streaming=True)
21test_segment = next(iter(ds_mcv_test))
22waveform = torch.from_numpy(test_segment["audio"]["array"])
23sample_rate = test_segment["audio"]["sampling_rate"]
24
25# Resample
26if sample_rate != model_sample_rate:
27 resampler = torchaudio.transforms.Resample(sample_rate, model_sample_rate)
28 waveform = resampler(waveform)
29
30# Get feat
31inputs = processor(waveform, sampling_rate=model_sample_rate, return_tensors="pt")
32input_features = inputs.input_features
33input_features = input_features.to(device)
34
35# Generate
36generated_ids = model.generate(inputs=input_features, max_new_tokens=225) # greedy
37# generated_ids = model.generate(inputs=input_features, max_new_tokens=225, num_beams=5) # beam search
38
39# Detokenize
40generated_sentences = processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
41
42# Normalise predicted sentences if necessary