Hugging Face 发布 Whisper 多语言 ASR 微调教程
Fine-Tune Whisper For Multilingual ASR with 🤗 Transformers
Hugging Face 发布使用 🤗 Transformers 对 Whisper 进行多语言语音识别(ASR)微调的分步教程。
完整给出从数据准备到训练和 Gradio 演示的代码流程,仅用 8 小时印地语数据把 WER 从 63.5% 降到 32%,方法可直接迁移到其他语言。
在这篇博客中,我们提供了一份分步指南,介绍如何使用 Hugging Face 🤗 Transformers 对 Whisper 进行微调,以适用于任何多语言 ASR 数据集。这篇博客深入解释了 Whisper 模型、Common Voice 数据集以及微调背后的理论,并附带了用于执行数据准备和微调步骤的代码单元。如需更精简的 notebook 版本(解释更少但包含所有代码),请参阅附带的 Google Colab。
目录
引言
Whisper 是一个用于自动语音识别(ASR)的预训练模型,由 OpenAI 的 Alec Radford 等人于 2022 年 9 月发布。与许多前身模型(如 Wav2Vec 2.0)不同,后者是在未标注的音频数据上进行预训练的,而 Whisper 是在大量标注的音频-转录数据上进行预训练的,准确地说,是 680,000 小时。这比用于训练 Wav2Vec 2.0 的未标注音频数据(60,000 小时)多了一个数量级。更重要的是,这些预训练数据中有 117,000 小时是多语言 ASR 数据。这使得其检查点可应用于超过 96 种语言,其中许多被认为是低资源语言。
如此数量的标注数据使得 Whisper 能够直接在语音识别的监督任务上进行预训练,从标注的音频-转录预训练数据中学习语音到文本的映射 1{}^1。因此,Whisper 只需很少的额外微调即可产生高性能的 ASR 模型。这与 Wav2Vec 2.0 形成对比,后者是在掩码预测的无监督任务上进行预训练的。在这里,模型仅从未标注的音频数据中学习语音到隐藏状态的中间映射。虽然无监督预训练能产生高质量的语音表示,但它不会学习语音到文本的映射。这种映射仅在微调期间学习,因此需要更多的微调才能产生有竞争力的性能。
当扩展到 680,000 小时的标注预训练数据时,Whisper 模型展现出强大的泛化能力,可适用于许多数据集和领域。预训练检查点取得了与最先进 ASR 系统相媲美的结果,在 LibriSpeech ASR 的 test-clean 子集上词错误率(WER)接近 3%,并在 TED-LIUM 上以 4.7% 的 WER 创下新的最先进水平(参见 Whisper 论文的表 8)。Whisper 在预训练期间获得的广泛多语言 ASR 知识可用于其他低资源语言;通过微调,预训练检查点可以适应特定的数据集和语言,以进一步改进这些结果。
Whisper 是一个基于 Transformer 的编码器-解码器模型,也被称为序列到序列模型。它将音频频谱图特征的序列映射到文本标记的序列。首先,原始音频输入通过特征提取器转换为对数梅尔频谱图。然后,Transformer 编码器对频谱图进行编码,形成编码器隐藏状态序列。最后,解码器自回归地预测文本标记,条件依赖于先前的标记和编码器隐藏状态。图 1 总结了 Whisper 模型。
在序列到序列模型中,编码器将音频输入转换为隐藏状态表示集,从语音中提取重要特征。解码器扮演语言模型的角色,处理隐藏状态表示并生成相应的文本转录。在系统架构中内部集成语言模型称为深度融合。这与浅融合形成对比,后者将语言模型与编码器外部结合,例如使用 CTC + nn-gram(参见 Internal Language Model Estimation)。通过深度融合,整个系统可以使用相同的训练数据和损失函数进行端到端训练,提供更大的灵活性并通常具有更优越的性能(参见 ESB Benchmark)。
Whisper 使用交叉熵目标函数进行预训练和微调,这是在分类任务上训练序列到序列系统的标准目标函数。在这里,系统被训练以从预定义的文本标记词汇表中正确分类目标文本标记。
Whisper 检查点有五种不同模型大小的配置。最小的四个在仅英语或多语言数据上训练。最大的检查点仅支持多语言。所有 11 个预训练检查点均可在 Hugging Face Hub 上获取。检查点在以下表格中总结,并附有 Hub 上模型的链接:
| 大小 | 层数 | 宽度 | 头数 | 参数 | 仅英语 | 多语言 |
|---|---|---|---|---|---|---|
| tiny | 4 | 384 | 6 | 39 M | ✓ | ✓ |
| base | 6 | 512 | 8 | 74 M | ✓ | ✓ |
| small | 12 | 768 | 12 | 244 M | ✓ | ✓ |
| medium | 24 | 1024 | 16 | 769 M | ✓ | ✓ |
| large | 32 | 1280 | 20 | 1550 M | x | ✓ |
| large-v2 | 32 | 1280 | 20 | 1550 M | x | ✓ |
| large-v3 | 32 | 1280 | 20 | 1550 M | x | ✓ |
为了演示目的,我们将使用 244M 参数(约 1GB)的 small 检查点的多语言版本进行微调。至于我们的数据,我们将在 Common Voice 数据集中的低资源语言上训练和评估我们的系统。我们将展示,仅用 8 小时的微调数据,我们就能在该语言上实现强大的性能。
1{}^1 Whisper 这个名字来源于缩写“WSPSR”,代表“Web-scale Supervised Pre-training for Speech Recognition”。
在 Google Colab 中微调 Whisper
准备环境
我们将使用几个流行的 Python 包来微调 Whisper 模型。
我们将使用 datasets[audio] 来下载和准备训练数据,同时使用
transformers 和 accelerate 来加载和训练我们的 Whisper 模型。
我们还需要 soundfile 包来预处理音频文件,
evaluate 和 jiwer 来评估我们模型的性能,以及
tensorboard 来记录我们的指标。最后,我们将使用 gradio 来构建一个
我们微调模型的炫酷演示。
!pip install --upgrade pip
!pip install --upgrade "datasets[audio]" transformers accelerate evaluate jiwer tensorboard gradio
我们强烈建议你在训练期间直接将模型检查点上传到 Hugging Face Hub。 Hub 提供:
- 集成的版本控制:你可以确保在训练过程中不会丢失任何模型检查点。
- Tensorboard 日志:在训练过程中跟踪重要指标。
- 模型卡片:记录模型的功能及其预期用例。
- 社区:一种与社区分享和协作的简便方式!
将 notebook 链接到 Hub 很简单——只需在提示时输入你的 Hub 身份验证令牌。在此处查找你的 Hub 身份验证令牌:
from huggingface_hub import notebook_login
notebook_login()
打印输出:
Login successful
Your token has been saved to /root/.huggingface/token
加载数据集
Common Voice 是一系列众包数据集,其中说话者 用各种语言录制来自维基百科的文本。我们将使用撰写本文时 最新版本的 Common Voice 数据集(版本 11)。 至于我们的语言,我们将在 印地语上微调我们的模型,这是一种印度-雅利安语, 在印度北部、中部、东部和西部使用。Common Voice 11.0 包含大约 12 小时的带标签印地语数据,其中 4 小时是 留出的测试数据。
提示:你可以通过查看 Hugging Face Hub 上的 Mozilla Foundation 组织页面来找到最新版本的 Common Voice 数据集。后续版本涵盖更多语言,并且每种语言包含更多数据。
让我们前往 Hub 并查看 Common Voice 的数据集页面:mozilla-foundation/common_voice_11_0。
第一次查看此页面时,我们会要求接受 使用条款。之后,我们将获得对数据集的完全访问权限。
一旦我们提供了使用数据集的身份验证,就会看到
数据集预览。数据集预览向我们展示了数据集的前 100 个样本。
更重要的是,它加载了音频样本,可供我们
实时收听。我们可以通过使用下拉菜单将子集设置为 hi 来选择 Common Voice 的印地语子集
(hi 是印地语的
语言标识符代码):
如果我们点击第一个样本上的播放按钮,就可以收听音频并 看到相应的文本。滚动浏览训练集 和测试集的样本,以更好地感受我们正在处理的音频和文本数据。你可以从语调和风格中看出,这些录音 取自旁白语音。你很可能还会注意到说话者和录音质量的巨大差异,这是众包数据的常见特征。
使用 🤗 Datasets,下载和准备数据极其简单。
我们只需一行代码即可下载和准备 Common Voice 拆分。
由于印地语资源非常少,我们将合并 train 和 validation
拆分,以提供大约 8 小时的训练数据。我们将使用 4 小时的
test 数据作为留出的测试集:
from datasets import load_dataset, DatasetDict
common_voice = DatasetDict()
common_voice["train"] = load_dataset("mozilla-foundation/common_voice_11_0", "hi", split="train+validation", use_auth_token=True)
common_voice["test"] = load_dataset("mozilla-foundation/common_voice_11_0", "hi", split="test", use_auth_token=True)
print(common_voice)
打印输出:
DatasetDict({
train: Dataset({
features: ['client_id', 'path', 'audio', 'sentence', 'up_votes', 'down_votes', 'age', 'gender', 'accent', 'locale', 'segment'],
num_rows: 6540
})
test: Dataset({
features: ['client_id', 'path', 'audio', 'sentence', 'up_votes', 'down_votes', 'age', 'gender', 'accent', 'locale', 'segment'],
num_rows: 2894
})
})
大多数 ASR 数据集仅提供输入音频样本(audio)和相应的转录文本(sentence)。Common Voice 包含额外的元数据信息,例如 accent 和 locale,对于 ASR 我们可以忽略这些信息。为了使 notebook 尽可能通用,我们仅考虑输入音频和转录文本进行微调,丢弃额外的元数据信息:
common_voice = common_voice.remove_columns(["accent", "age", "client_id", "down_votes", "gender", "locale", "path", "segment", "up_votes"])
Common Voice 只是我们可以从 Hub 下载的多语言 ASR 数据集之一——还有更多可供我们使用!要查看可用于语音识别的数据集范围,请点击链接:Hub 上的 ASR 数据集。
准备特征提取器、分词器和数据
ASR 流水线可以分解为三个组件:
- 一个特征提取器,用于预处理原始音频输入
- 执行序列到序列映射的模型
- 一个分词器,用于将模型输出后处理为文本格式
在 🤗 Transformers 中,Whisper 模型有一个关联的特征提取器和分词器,分别称为 WhisperFeatureExtractor 和 WhisperTokenizer。
我们将逐一介绍特征提取器和分词器的细节!
加载 WhisperFeatureExtractor
语音由随时间变化的一维数组表示。数组在任何给定时间步的值是该点信号的振幅。仅从振幅信息,我们就可以重建音频的频谱并恢复所有声学特征。
由于语音是连续的,它包含无限数量的振幅值。这给期望有限数组的计算机设备带来了问题。因此,我们通过在固定时间步从信号中采样值来离散化语音信号。我们采样音频的间隔称为采样率,通常以样本/秒或赫兹 (Hz) 为单位。以更高的采样率进行采样可以更好地近似连续语音信号,但也需要每秒存储更多的值。
将音频输入的采样率与我们模型期望的采样率匹配至关重要,因为不同采样率的音频信号具有非常不同的分布。音频样本只应以正确的采样率进行处理。否则可能导致意外结果!例如,以 16kHz 的采样率获取音频样本并以 8kHz 的采样率收听它,会使音频听起来像是半速。同样,传递错误采样率的音频可能会使期望一种采样率却收到另一种采样率的 ASR 模型出错。Whisper 特征提取器期望音频输入的采样率为 16kHz,因此我们需要将输入匹配到这个值。我们不想无意中在慢动作语音上训练 ASR 系统!
Whisper 特征提取器执行两个操作。它首先对一批音频样本进行填充/截断,使所有样本的输入长度都为 30 秒。短于 30 秒的样本通过在序列末尾追加零来填充到 30 秒(音频信号中的零对应于无信号或静音)。长于 30 秒的样本被截断为 30 秒。由于批次中的所有元素在输入空间中都被填充/截断到最大长度,因此在将音频输入前向传播到 Whisper 模型时,我们不需要注意力掩码。Whisper 在这方面是独特的——对于大多数音频模型,你需要提供一个注意力掩码,详细说明序列在哪里被填充,从而在自注意力机制中应忽略哪些位置。Whisper 被训练为在没有注意力掩码的情况下运行,并直接从语音信号中推断出在哪里忽略输入。
Whisper 特征提取器执行的第二个操作是将填充后的音频数组转换为对数梅尔频谱图。这些频谱图是信号频率的可视化表示,类似于傅里叶变换。图 2 显示了一个示例频谱图。沿 yy 轴是梅尔通道,对应于特定的频率区间。沿 xx 轴是时间。每个像素的颜色对应于给定时间该频率区间的对数强度。对数梅尔频谱图是 Whisper 模型期望的输入形式。
梅尔通道(频率区间)在语音处理中是标准的,并被选择来近似人类的听觉范围。对于 Whisper 微调,我们只需要知道频谱图是语音信号中频率的可视化表示。有关梅尔通道的更多细节,请参阅 梅尔频率倒谱。
幸运的是,🤗 Transformers 的 Whisper 特征提取器只需一行代码即可完成填充和频谱图转换!让我们继续从预训练检查点加载特征提取器,为我们的音频数据做好准备:
from transformers import WhisperFeatureExtractor
feature_extractor = WhisperFeatureExtractor.from_pretrained("openai/whisper-small")
加载 WhisperTokenizer
现在让我们看看如何加载 Whisper 分词器。Whisper 模型输出文本标记,这些标记指示预测文本在词汇项字典中的索引。分词器将文本标记序列映射到实际的文本字符串(例如 [1169, 3797, 3332] -> "the cat sat")。
传统上,当使用仅编码器模型进行 ASR 时,我们使用 连接主义时间分类(CTC) 进行解码。这里我们需要为我们使用的每个数据集训练一个 CTC 分词器。使用编码器-解码器架构的优势之一是我们可以直接利用预训练模型中的分词器。
Whisper 分词器是在 96 种预训练语言的转录文本上预训练的。 因此,它拥有一个广泛的 字节对, 适用于几乎所有多语言 ASR 应用。 对于印地语,我们可以加载分词器并将其用于微调,而无需 任何进一步的修改。我们只需指定目标语言 和任务。这些参数会告知分词器在编码标签序列的开头添加语言 和任务标记:
from transformers import WhisperTokenizer
tokenizer = WhisperTokenizer.from_pretrained("openai/whisper-small", language="Hindi", task="transcribe")
提示:通过在上述行中将任务设置为
"translate"并将语言设置为目标文本语言,这篇博客文章可以改编用于语音翻译。这样在预处理数据集时,会为语音翻译添加相关的任务和语言标记。
我们可以通过编码和解码 Common Voice 数据集的第一个样本,来验证分词器是否正确编码了印地语字符。在 编码转录文本时,分词器会在序列的开头和结尾附加“特殊标记”,包括转录文本的开始/结束标记、 语言标记和任务标记(如 上一步中的参数所指定)。 在解码标签 id 时,我们可以选择“跳过”这些特殊 标记,从而以原始输入形式返回字符串:
input_str = common_voice["train"][0]["sentence"]
labels = tokenizer(input_str).input_ids
decoded_with_special = tokenizer.decode(labels, skip_special_tokens=False)
decoded_str = tokenizer.decode(labels, skip_special_tokens=True)
print(f"Input: {input_str}")
print(f"Decoded w/ special: {decoded_with_special}")
print(f"Decoded w/out special: {decoded_str}")
print(f"Are equal: {input_str == decoded_str}")
打印输出:
Input: खीर की मिठास पर गरमाई बिहार की सियासत, कुशवाहा ने दी सफाई
Decoded w/ special: <|startoftranscript|><|hi|><|transcribe|><|notimestamps|>खीर की मिठास पर गरमाई बिहार की सियासत, कुशवाहा ने दी सफाई<|endoftext|>
Decoded w/out special: खीर की मिठास पर गरमाई बिहार की सियासत, कुशवाहा ने दी सफाई
Are equal: True
组合创建 WhisperProcessor
为了简化特征提取器和分词器的使用,我们可以将
两者包装到一个单一的 WhisperProcessor 类中。这个处理器对象
继承自 WhisperFeatureExtractor 和 WhisperProcessor,
并可根据需要在音频输入和模型预测上使用。
这样一来,我们在训练期间只需跟踪两个对象:
processor 和 model:
from transformers import WhisperProcessor
processor = WhisperProcessor.from_pretrained("openai/whisper-small", language="Hindi", task="transcribe")
准备数据
让我们打印 Common Voice 数据集的第一个示例,看看 数据是什么形式:
print(common_voice["train"][0])
打印输出:
{'audio': {'path': '/home/sanchit_huggingface_co/.cache/huggingface/datasets/downloads/extracted/607848c7e74a89a3b5225c0fa5ffb9470e39b7f11112db614962076a847f3abf/cv-corpus-11.0-2022-09-21/hi/clips/common_voice_hi_25998259.mp3',
'array': array([0.0000000e+00, 0.0000000e+00, 0.0000000e+00, ..., 9.6724887e-07,
1.5334779e-06, 1.0415988e-06], dtype=float32),
'sampling_rate': 48000},
'sentence': 'खीर की मिठास पर गरमाई बिहार की सियासत, कुशवाहा ने दी सफाई'}
我们可以看到,我们得到了一个一维输入音频数组以及 对应的目标转录文本。我们已经详细讨论过 采样率的重要性,以及我们需要将音频的 采样率与 Whisper 模型(16kHz)相匹配。由于 我们的输入音频采样率为 48kHz,我们需要在将其传递给 Whisper 特征提取器之前将其下采样到 16kHz。
我们将使用数据集的
cast_column
方法将音频输入设置为正确的采样率。此操作不会就地更改音频,
而是向 datasets 发出信号,使其在首次加载音频样本时即时重采样:
from datasets import Audio
common_voice = common_voice.cast_column("audio", Audio(sampling_rate=16000))
重新加载 Common Voice 数据集中的第一个音频样本,会将其重采样 到所需的采样率:
print(common_voice["train"][0])
打印输出:
{'audio': {'path': '/home/sanchit_huggingface_co/.cache/huggingface/datasets/downloads/extracted/607848c7e74a89a3b5225c0fa5ffb9470e39b7f11112db614962076a847f3abf/cv-corpus-11.0-2022-09-21/hi/clips/common_voice_hi_25998259.mp3',
'array': array([ 0.0000000e+00, 0.0000000e+00, 0.0000000e+00, ...,
-3.4206650e-07, 3.2979898e-07, 1.0042874e-06], dtype=float32),
'sampling_rate': 16000},
'sentence': 'खीर की मिठास पर गरमाई बिहार की सियासत, कुशवाहा ने दी सफाई'}
很好!我们可以看到采样率已下采样到 16kHz。数组 值也不同了,因为现在大约每三个先前的振幅值 只对应一个振幅值。
现在我们可以编写一个函数来准备数据,使其可供模型使用:
- 我们通过调用
batch["audio"]来加载并重采样音频数据。如上所述,🤗 Datasets 会即时执行任何必要的重采样操作。 - 我们使用特征提取器从一维音频数组计算对数梅尔频谱图输入特征。
- 我们通过使用分词器将转录文本编码为标签 id。
def prepare_dataset(batch):
# load and resample audio data from 48 to 16kHz
audio = batch["audio"]
# compute log-Mel input features from input audio array
batch["input_features"] = feature_extractor(audio["array"], sampling_rate=audio["sampling_rate"]).input_features[0]
# encode target text to label ids
batch["labels"] = tokenizer(batch["sentence"]).input_ids
return batch
我们可以使用数据集的 .map 方法将数据准备函数应用于所有训练示例:
common_voice = common_voice.map(prepare_dataset, remove_columns=common_voice.column_names["train"], num_proc=4)
好的!这样我们就为训练完全准备好了数据! 让我们继续,看看如何使用这些数据来微调 Whisper。
注意:目前 datasets 同时使用 torchaudio 和 librosa 来进行音频加载和重采样。
如果你想实现自定义的数据加载/采样,可以使用 "path" 列来获取音频文件路径,并忽略 "audio" 列。
训练与评估
现在我们已经准备好了数据,可以深入训练流程了。 🤗 Trainer 将为我们完成大部分繁重的工作。我们只需要:
加载预训练检查点:我们需要加载一个预训练检查点,并为其正确配置训练。
定义数据整理器:数据整理器接收我们预处理后的数据,并准备好可供模型使用的 PyTorch 张量。
评估指标:在评估过程中,我们希望使用 词错误率(WER) 指标来评估模型。我们需要定义一个
compute_metrics函数来处理该计算。定义训练参数:这些参数将由 🤗 Trainer 用于构建训练计划。
一旦我们微调好模型,我们将在测试数据上对其进行评估,以验证我们已正确训练它 来转写印地语语音。
加载预训练检查点
我们将从预训练的 Whisper small 检查点开始微调运行。
为此,我们将从 Hugging Face Hub 加载预训练权重。
同样,通过使用 🤗 Transformers,这非常简单!
from transformers import WhisperForConditionalGeneration
model = WhisperForConditionalGeneration.from_pretrained("openai/whisper-small")
在推理时,Whisper 模型会自动检测源音频的语言,
并预测该语言的 token id。
在源音频语言先验已知的情况下,例如
多语言微调,显式设置语言是有益的。
这可以避免预测出错误语言的情况,
从而导致生成过程中预测文本偏离真实语言。为此,我们将
生成配置中的 language
和 task
参数设置为相应值。我们还会将任何 forced_decoder_ids
设置为 None,因为这是设置语言和任务参数的旧方式:
model.generation_config.language = "hindi"
model.generation_config.task = "transcribe"
model.generation_config.forced_decoder_ids = None
定义数据整理器
序列到序列语音模型的数据整理器的独特之处在于,
它独立处理 input_features 和 labels:input_features 必须
由特征提取器处理,而 labels 由分词器处理。
input_features 已经填充到 30 秒并转换为固定维度的
对数梅尔频谱图,所以我们只需将它们转换为批处理的 PyTorch 张量。我们使用
特征提取器的 .pad 方法并传入 return_tensors=pt 来完成此操作。注意,这里不会应用额外的
填充,因为输入是固定维度的,
input_features 只是被转换为 PyTorch 张量。
另一方面,labels 未填充。我们首先使用分词器的 .pad 方法
将序列填充到批次中的最大长度。然后,填充 token
会被替换为 -100,以便在
计算损失时不考虑这些 token。接着,我们从标签序列开头切掉起始转写 token,因为
我们会在训练期间稍后追加它。
我们可以利用之前定义的 WhisperProcessor 来执行
特征提取器和分词器操作:
import torch
from dataclasses import dataclass
from typing import Any, Dict, List, Union
@dataclass
class DataCollatorSpeechSeq2SeqWithPadding:
processor: Any
decoder_start_token_id: int
def __call__(self, features: List[Dict[str, Union[List[int], torch.Tensor]]]) -> Dict[str, torch.Tensor]:
# split inputs and labels since they have to be of different lengths and need different padding methods
# first treat the audio inputs by simply returning torch tensors
input_features = [{"input_features": feature["input_features"]} for feature in features]
batch = self.processor.feature_extractor.pad(input_features, return_tensors="pt")
# get the tokenized label sequences
label_features = [{"input_ids": feature["labels"]} for feature in features]
# pad the labels to max length
labels_batch = self.processor.tokenizer.pad(label_features, return_tensors="pt")
# replace padding with -100 to ignore loss correctly
labels = labels_batch["input_ids"].masked_fill(labels_batch.attention_mask.ne(1), -100)
# if bos token is appended in previous tokenization step,
# cut bos token here as it's append later anyways
if (labels[:, 0] == self.decoder_start_token_id).all().cpu().item():
labels = labels[:, 1:]
batch["labels"] = labels
return batch
让我们初始化刚刚定义的数据整理器:
data_collator = DataCollatorSpeechSeq2SeqWithPadding(
processor=processor,
decoder_start_token_id=model.config.decoder_start_token_id,
)
评估指标
接下来,我们定义将在评估集上使用的评估指标。我们将使用词错误率(WER)指标,这是评估 ASR 系统的“事实上的”指标。有关更多信息,请参阅 WER 文档。我们将从 🤗 Evaluate 加载 WER 指标:
import evaluate
metric = evaluate.load("wer")
然后,我们只需定义一个函数,该函数接收我们的模型预测并返回 WER 指标。这个函数称为 compute_metrics,它首先将 label_ids 中的 -100 替换为 pad_token_id(撤销我们在数据整理器中应用的操作,以便在损失中正确忽略填充标记)。然后,它将预测和标签 id 解码为字符串。最后,它计算预测与参考标签之间的 WER:
def compute_metrics(pred):
pred_ids = pred.predictions
label_ids = pred.label_ids
# replace -100 with the pad_token_id
label_ids[label_ids == -100] = tokenizer.pad_token_id
# we do not want to group tokens when computing the metrics
pred_str = tokenizer.batch_decode(pred_ids, skip_special_tokens=True)
label_str = tokenizer.batch_decode(label_ids, skip_special_tokens=True)
wer = 100 * metric.compute(predictions=pred_str, references=label_str)
return {"wer": wer}
定义训练参数
在最后一步中,我们定义与训练相关的所有参数。下面解释了其中一部分参数:
output_dir:用于保存模型权重的本地目录。这也将是 Hugging Face Hub 上的仓库名称。generation_max_length:评估期间自回归生成的最大 token 数。save_steps:在训练期间,中间检查点将每save_steps个训练步骤异步保存并上传到 Hub。eval_steps:在训练期间,将每eval_steps个训练步骤对中间检查点进行一次评估。report_to:保存训练日志的位置。支持的平台有"azure_ml"、"comet_ml"、"mlflow"、"neptune"、"tensorboard"和"wandb"。选择你喜欢的平台,或保留为"tensorboard"以记录到 Hub。
有关其他训练参数的更多详细信息,请参阅 Seq2SeqTrainingArguments 文档。
from transformers import Seq2SeqTrainingArguments
training_args = Seq2SeqTrainingArguments(
output_dir="./whisper-small-hi", # change to a repo name of your choice
per_device_train_batch_size=16,
gradient_accumulation_steps=1, # increase by 2x for every 2x decrease in batch size
learning_rate=1e-5,
warmup_steps=500,
max_steps=5000,
gradient_checkpointing=True,
fp16=True,
evaluation_strategy="steps",
per_device_eval_batch_size=8,
predict_with_generate=True,
generation_max_length=225,
save_steps=1000,
eval_steps=1000,
logging_steps=25,
report_to=["tensorboard"],
load_best_model_at_end=True,
metric_for_best_model="wer",
greater_is_better=False,
push_to_hub=True,
)
注意:如果不想将模型检查点上传到 Hub,请设置 push_to_hub=False。
我们可以将训练参数连同我们的模型、数据集、数据整理器和 compute_metrics 函数一起转发给 🤗 Trainer:
from transformers import Seq2SeqTrainer
trainer = Seq2SeqTrainer(
args=training_args,
model=model,
train_dataset=common_voice["train"],
eval_dataset=common_voice["test"],
data_collator=data_collator,
compute_metrics=compute_metrics,
tokenizer=processor.feature_extractor,
)
至此,我们已准备好开始训练!
训练
要启动训练,只需执行:
trainer.train()
训练大约需要 5-10 小时,具体取决于你的 GPU 或分配给 Google Colab 的 GPU。根据你的 GPU 情况,你在开始训练时可能会遇到 CUDA "out-of-memory" 错误。在这种情况下,你可以将 per_device_train_batch_size 逐步减半,并使用 gradient_accumulation_steps 来补偿。
打印输出:
| Step | Training Loss | Epoch | Validation Loss | WER |
|---|---|---|---|---|
| 1000 | 0.1011 | 2.44 | 0.3075 | 34.63 |
| 2000 | 0.0264 | 4.89 | 0.3558 | 33.13 |
| 3000 | 0.0025 | 7.33 | 0.4214 | 32.59 |
| 4000 | 0.0006 | 9.78 | 0.4519 | 32.01 |
| 5000 | 0.0002 | 12.22 | 0.4679 | 32.10 |
经过 4000 个训练步骤后,我们最佳的 WER 为 32.0%。作为参考,预训练的 Whisper small 模型达到了 63.5% 的 WER,这意味着我们通过微调实现了 31.5% 的绝对提升。仅用 8 小时训练数据就有这样的效果,还不错!
我们现在准备在 Hugging Face Hub 上分享我们微调后的模型。为了通过适当的标签和 README 信息使其更易于访问,我们可以在推送时设置适当的关键字参数(kwargs)。你可以相应地更改这些值,以匹配你的数据集、语言和模型名称:
kwargs = {
"dataset_tags": "mozilla-foundation/common_voice_11_0",
"dataset": "Common Voice 11.0", # a 'pretty' name for the training dataset
"dataset_args": "config: hi, split: test",
"language": "hi",
"model_name": "Whisper Small Hi - Sanchit Gandhi", # a 'pretty' name for your model
"finetuned_from": "openai/whisper-small",
"tasks": "automatic-speech-recognition",
}
现在可以将训练结果上传到 Hub。为此,请执行 push_to_hub 命令:
trainer.push_to_hub(**kwargs)
现在,你可以使用 Hub 上的链接与任何人分享此模型。他们还可以使用标识符 "your-username/the-name-you-picked" 加载它,例如:
from transformers import WhisperForConditionalGeneration, WhisperProcessor
model = WhisperForConditionalGeneration.from_pretrained("sanchit-gandhi/whisper-small-hi")
processor = WhisperProcessor.from_pretrained("sanchit-gandhi/whisper-small-hi")
虽然微调后的模型在 Common Voice Hindi 测试数据上产生了令人满意的结果,但它绝不是最优的。本 notebook 的目的是演示如何在任何多语言 ASR 数据集上微调预训练的 Whisper 检查点。通过优化训练超参数(例如学习率和dropout),并使用更大的预训练检查点(medium 或 large-v3),结果可能会得到改善。
构建演示
现在我们已经微调好了模型,可以构建一个演示来展示它的 ASR 能力!我们将使用 🤗 Transformers pipeline,它会处理整个 ASR 流程,从预处理音频输入到解码模型预测。我们将使用 Gradio 构建交互式演示。Gradio 可以说是构建机器学习演示最直接的方式;有了 Gradio,我们只需几分钟就能构建一个演示!
运行下面的示例将生成一个 Gradio 演示,我们可以通过计算机的麦克风录制语音,并将其输入到我们微调好的 Whisper 模型中以转录相应的文本:
from transformers import pipeline
import gradio as gr
pipe = pipeline(model="sanchit-gandhi/whisper-small-hi") # change to "your-username/the-name-you-picked"
def transcribe(audio):
text = pipe(audio)["text"]
return text
iface = gr.Interface(
fn=transcribe,
inputs=gr.Audio(source="microphone", type="filepath"),
outputs="text",
title="Whisper Small Hindi",
description="Realtime demo for Hindi speech recognition using a fine-tuned Whisper small model.",
)
iface.launch()
结语
在这篇博客中,我们介绍了使用 🤗 Datasets、Transformers 和 Hugging Face Hub 为多语言 ASR 微调 Whisper 的分步指南。如果你想自己尝试微调,请参考 Google Colab。如果你有兴趣微调其他 Transformers 模型,无论是用于英语还是多语言 ASR,请务必查看 examples/pytorch/speech-recognition 中的示例脚本。
来源:Hugging Face:Blog · huggingface.co