Views
No views yet
import torchaudio
from datasets import load_dataset, load_metric
from transformers import (
Wav2Vec2ForCTC,
Wav2Vec2Processor,
)
import torch
import re
import sys
model_name = "voidful/wav2vec2-large-xlsr-53-hk"
device = "cuda"
processor_name = "voidful/wav2vec2-large-xlsr-53-hk"
chars_to_ignore_regex = r"[¥•"#$%&'()*+,-/:;<=>@[\]^_`{|}~⦅⦆「」、 、〃〈〉《》「」『』【】〔〕〖〗〘〙〚〛〜〝〞〟〰〾〿–—‘’‛“”„‟…‧﹏﹑﹔·'℃°•·.﹑︰〈〉─《﹖﹣﹂﹁﹔!?。。"#$%&'()*+,﹐-/:;<=>@[\]^_`{|}~⦅⦆「」、、〃》「」『』【】〔〕〖〗〘〙〚〛〜〝〞〟〰〾〿–—‘’‛“”„‟…‧﹏..!\\"#$%&()*+,\\-.\\:;<=>?@\\[\\]\\\\\\/^_`{|}~]"
model = Wav2Vec2ForCTC.from_pretrained(model_name).to(device)
processor = Wav2Vec2Processor.from_pretrained(processor_name)
resampler = torchaudio.transforms.Resample(orig_freq=48_000, new_freq=16_000)
def load_file_to_data(file):
batch = {}
speech, _ = torchaudio.load(file)
batch["speech"] = resampler.forward(speech.squeeze(0)).numpy()
batch["sampling_rate"] = resampler.new_freq
return batch
def predict(data):
features = processor(data["speech"], sampling_rate=data["sampling_rate"], padding=True, return_tensors="pt")
input_values = features.input_values.to(device)
attention_mask = features.attention_mask.to(device)
with torch.no_grad():
logits = model(input_values, attention_mask=attention_mask).logits
pred_ids = torch.argmax(logits, dim=-1)
return processor.batch_decode(pred_ids)
predict(load_file_to_data('voice file path'))1!mkdir cer
2!wget -O cer/cer.py https://huggingface.co/ctl/wav2vec2-large-xlsr-cantonese/raw/main/cer.py
3!pip install jiwer
4
5import torchaudio
6from datasets import load_dataset, load_metric
7from transformers import (
8 Wav2Vec2ForCTC,
9 Wav2Vec2Processor,
10)
11import torch
12import re
13import sys
14
15cer = load_metric("./cer")
16model_name = "voidful/wav2vec2-large-xlsr-53-hk"
17device = "cuda"
18processor_name = "voidful/wav2vec2-large-xlsr-53-hk"
19
20chars_to_ignore_regex = r"[¥•"#$%&'()*+,-/:;<=>@[\]^_`{|}~⦅⦆「」、 、〃〈〉《》「」『』【】〔〕〖〗〘〙〚〛〜〝〞〟〰〾〿–—‘’‛“”„‟…‧﹏﹑﹔·'℃°•·.﹑︰〈〉─《﹖﹣﹂﹁﹔!?。。"#$%&'()*+,﹐-/:;<=>@[\]^_`{|}~⦅⦆「」、、〃》「」『』【】〔〕〖〗〘〙〚〛〜〝〞〟〰〾〿–—‘’‛“”„‟…‧﹏..!\\"#$%&()*+,\\-.\\:;<=>?@\\[\\]\\\\\\/^_`{|}~]"
21
22model = Wav2Vec2ForCTC.from_pretrained(model_name).to(device)
23processor = Wav2Vec2Processor.from_pretrained(processor_name)
24
25ds = load_dataset("common_voice", 'zh-HK', data_dir="./cv-corpus-6.1-2020-12-11", split="test")
26
27resampler = torchaudio.transforms.Resample(orig_freq=48_000, new_freq=16_000)
28
29def map_to_array(batch):
30 speech, _ = torchaudio.load(batch["path"])
31 batch["speech"] = resampler.forward(speech.squeeze(0)).numpy()
32 batch["sampling_rate"] = resampler.new_freq
33 batch["sentence"] = re.sub(chars_to_ignore_regex, '', batch["sentence"]).lower().replace("’", "'")
34 return batch
35
36ds = ds.map(map_to_array)
37
38def map_to_pred(batch):
39 features = processor(batch["speech"], sampling_rate=batch["sampling_rate"][0], padding=True, return_tensors="pt")
40 input_values = features.input_values.to(device)
41 attention_mask = features.attention_mask.to(device)
42 with torch.no_grad():
43 logits = model(input_values, attention_mask=attention_mask).logits
44 pred_ids = torch.argmax(logits, dim=-1)
45 batch["predicted"] = processor.batch_decode(pred_ids)
46 batch["target"] = batch["sentence"]
47 return batch
48
49result = ds.map(map_to_pred, batched=True, batch_size=16, remove_columns=list(ds.features.keys()))
50
51print("CER: {:2f}".format(100 * cer.compute(predictions=result["predicted"], references=result["target"])))CER 16.41