跳到正文
原文
Hugging Face:Blog(RSS)·· 2023-12-20精选AI 评分60

Hugging Face 讲解用 Speculative Decoding 将 Whisper 推理提速 2 倍

Speculative Decoding for 2x Faster Whisper Inference

AI 导读

Hugging Face 发布教程,演示用 Speculative Decoding 将 Whisper 推理提速约 2 倍,且数学上保证输出与主模型单独推理完全一致,可作为现有 Whisper 流程的即插即用替换。

推荐理由

原文给出可复现的 Whisper 推理提速方法和实测数字,读者可以直接迁移到自己的语音转写流程。

正文 · AI 翻译

Open In Colab

Open AI 的 Whisper 是一个通用的语音转录模型,在一系列不同的基准测试和音频条件下都取得了最先进的结果。最新的 large-v3 模型在 OpenASR Leaderboard 上名列前茅,被评为最佳的英语开源语音转录模型。该模型还展现出强大的多语言性能,在 Common Voice 15 数据集测试的 58 种语言中,有 42 种语言的词错误率(WER)低于 30%。

虽然转录准确率非常出色,但推理时间非常慢。一段 1 小时的音频片段在 16GB T4 GPU 上转录需要超过 6 分钟,即使利用了 flash attention、半精度和 chunking 等推理优化。

在这篇博客文章中,我们展示了如何使用投机解码将 Whisper 的推理时间减少 2 倍,同时在数学上确保模型产生完全相同的输出。因此,该方法可以完美地直接替换现有的 Whisper 流水线,因为它在保持相同准确率的同时提供了免费的 2 倍加速。如需更精简的博客文章版本(解释更少但包含所有代码),请参阅随附的 Google Colab。

投机解码

投机解码由 Google 的 Yaniv Leviathan 等人在 Fast Inference from Transformers via Speculative Decoding 中提出。它的工作原理基于一个前提:更快的助手模型经常生成与更大的主模型相同的 token。

首先,助手模型自回归地生成一个包含 N N 个候选 token 的序列,y^1:N \hat{\boldsymbol{y}}_{1:N} 。在下图中,助手模型生成了一个包含 5 个候选 token 的序列:The quick brown sock jumps。

虽然这些候选 token 生成得很快,但它们可能与主模型预测的 token 不同。因此,在第二步中,候选 token 被传递给主模型进行“验证”。主模型将候选 token 作为输入,执行单次前向传播。主模型的输出是 token 序列中每一步的“正确”token y1:N \boldsymbol{y}_{1:N} 。

在上图中,我们看到主模型预测的前三个 token 与助手模型的预测一致:The quick brown。然而,助手模型的第四个候选 token sock 与主模型的正确 token fox 不匹配。

我们知道,在第一个不匹配之前的所有候选 token 都是正确的(The quick brown),因为这些与主模型的预测一致。然而,在第一个不匹配之后,候选 token 偏离了主模型预测的实际 token。因此,我们可以将第一个不正确的候选 token(sock)替换为主模型的正确 token(fox),并丢弃此后所有预测的 token,因为这些已经偏离了。修正后的序列 The quick brown fox 现在构成了助手模型的新输入:

随后推理过程重复进行,助手模型生成一组新的 N N 个候选 token,由主模型在单次前向传播中验证。

由于我们使用快速的助手模型进行自回归生成,并且仅使用缓慢的主模型执行验证前向传播,解码过程被大幅加速。此外,主模型执行的验证前向传播确保获得完全相同的输出,就如同我们单独使用主模型一样。这使得推测解码可以完美地直接替换现有的 Whisper 流水线,因为可以确信将达到相同的质量。

为了获得最大的延迟改进,助手模型应显著快于主模型,同时尽可能频繁地预测相同的 token 分布。在实践中,这两个属性形成一种权衡:模型越快,其准确性越低。然而,由于所有预测 token 中 70-80% 往往是“较容易”的 token,这种权衡严重偏向于选择更快的模型,而不是更准确的模型。因此,助手模型应至少比主模型快 3 倍(越快越好),同时正确预测示例中所有“容易”的 token。剩余的 20-30% 更“困难”的 token 则可以由更大的主模型来验证。

选择助手模型的唯一约束是它必须与主模型共享相同的词汇表。也就是说,助手模型必须使用与主模型一对一的相同分词器。因此,如果我们想对 Whisper 的多语言变体(例如 large-v2(多语言))使用推测解码,我们需要选择 Whisper 的多语言变体作为助手模型,例如 tiny。而如果我们想对仅英语版本的 Whisper(例如 medium.en)使用推测解码,我们需要一个仅英语版本作为助手模型,例如 tiny.en。目前,Whisper large-v3 是一个例外,因为它是唯一一个具有扩展词汇表大小的 Whisper 检查点,因此与之前的 Whisper 检查点不兼容。

现在我们已经了解了推测解码背后的背景,准备深入实际实现。在 🤗 Transformers 库中,推测解码被实现为“辅助生成”推理策略。有关实现的更多细节,建议读者阅读 Joao Gante 关于辅助生成的精彩博客文章。

英语语音转录

基线实现

我们首先对 Whisper large-v2 进行基准测试,以获得推理速度的基线数值。我们可以通过便捷的 AutoModelForSpeechSeq2Seq 和 AutoProcessor 类加载主模型及其对应的处理器。我们将以 float16 精度加载模型,并通过传递 low_cpu_mem_usage=True 确保加载时间尽可能短。此外,我们希望通过传递 use_safetensors=True 确保模型以 safetensors 格式加载。最后,我们将传递参数 attn_implementation="sdpa",以通过 PyTorch 的 SDPA 注意力内核获得 Flash Attention 加速:

import torch
from transformers import AutoModelForSpeechSeq2Seq, AutoProcessor

device = "cuda:0" if torch.cuda.is_available() else "cpu"
torch_dtype = torch.float16 if torch.cuda.is_available() else torch.float32

model_id = "openai/whisper-large-v2"

model = AutoModelForSpeechSeq2Seq.from_pretrained(
    model_id,
    torch_dtype=torch_dtype,
    low_cpu_mem_usage=True,
    use_safetensors=True,
    attn_implementation="sdpa",
)
model.to(device)

processor = AutoProcessor.from_pretrained(model_id)

让我们加载将用于基准测试的英语语音转录数据集。我们将加载一个由 LibriSpeech ASR 验证干净数据集中的 73 个样本组成的小型数据集。这大约相当于 9MB 的数据,因此非常轻量,在设备上可以快速下载:

from datasets import load_dataset

dataset = load_dataset("hf-internal-testing/librispeech_asr_dummy", "clean", split="validation")

对于基准测试,我们只想测量生成时间,所以让我们编写一个简短的辅助函数来测量这一步。以下函数将返回解码后的 token 以及运行模型所花费的时间:

import time

def generate_with_time(model, inputs, **kwargs):
    start_time = time.time()
    outputs = model.generate(**inputs, **kwargs)
    generation_time = time.time() - start_time
    return outputs, generation_time

现在我们可以遍历数据集中的音频样本,并累加总体生成时间:

from tqdm import tqdm

all_time = 0
predictions = []
references = []

for sample in tqdm(dataset):
    audio = sample["audio"]
    inputs = processor(audio["array"], sampling_rate=audio["sampling_rate"], return_tensors="pt")
    inputs = inputs.to(device=device, dtype=torch.float16)
    
    output, gen_time = generate_with_time(model, inputs)
    all_time += gen_time
    predictions.append(processor.batch_decode(output, skip_special_tokens=True, normalize=True)[0])
    references.append(processor.tokenizer._normalize(sample["text"]))

print(all_time)

输出:

100%|██████████| 73/73 [01:37<00:00,  1.33s/it]
72.99542546272278

好的!我们看到转录这 73 个样本花费了 73 秒。让我们检查预测的 WER:

from evaluate import load

wer = load("wer")
print(wer.compute(predictions=predictions, references=references))

输出:

0.03507271171941831

我们的最终基线数字是 73 秒,WER 为 3.5%。

投机解码

现在让我们加载用于投机解码的助手模型。在这个例子中,我们将使用 Whisper 的蒸馏变体,distil-large-v2。蒸馏模型复制了 Whisper 的整个编码器,但只复制了 32 个解码器层中的 2 个。因此,它的运行速度比 Whisper 快 6 倍,同时在分布外测试集上的 WER 表现差距在 1% 以内。这使其成为助手模型的完美选择,因为它既具有高转录准确率,又具有快速生成能力 1{}^1。

由于 Distil-Whisper 使用与 Whisper 模型完全相同的编码器,我们可以在主模型和助手模型之间共享编码器。然后,我们只需将 Distil-Whisper 的 2 层解码器作为“仅解码器”模型加载。我们可以通过方便的 AutoModelForCausalLM 自动类来实现这一点。在实践中,与单独使用主模型相比,这仅导致 VRAM 增加 8%。

from transformers import AutoModelForCausalLM

assistant_model_id = "distil-whisper/distil-large-v2"

assistant_model = AutoModelForCausalLM.from_pretrained(
    assistant_model_id,
    torch_dtype=torch_dtype,
    low_cpu_mem_usage=True,
    use_safetensors=True,
    attn_implementation="sdpa",
)

assistant_model.to(device)

1{}^1 我们打算发布一个改进版的 Distil-Whisper,它在 token 分布上具有更强的一致性,这将进一步提高投机解码的性能。请关注 Distil-Whisper 仓库 以获取更新。


我们可以为投机解码基准测试定义一个修改后的函数。与之前函数的唯一区别是,我们在调用 .generate 时传入了助手模型:

def assisted_generate_with_time(model, inputs, **kwargs):
    start_time = time.time()
    outputs = model.generate(**inputs, assistant_model=assistant_model, **kwargs)
    generation_time = time.time() - start_time
    return outputs, generation_time

让我们使用 Distil-Whisper 作为 Whisper 的助手来运行投机解码基准测试:

all_time = 0
predictions = []
references = []

for sample in tqdm(dataset):
    audio = sample["audio"]
    inputs = processor(audio["array"], sampling_rate=audio["sampling_rate"], return_tensors="pt")
    inputs = inputs.to(device=device, dtype=torch.float16)
    
    output, gen_time = assisted_generate_with_time(model, inputs)
    all_time += gen_time
    predictions.append(processor.batch_decode(output, skip_special_tokens=True, normalize=True)[0])
    references.append(processor.tokenizer._normalize(sample["text"]))

print(all_time)

输出:

100%|██████████| 73/73 [00:38<00:00,  1.88it/s]
32.69683289527893

使用投机解码,推理时间仅为 33 秒,比之前快了 2.2 倍!让我们验证一下是否具有相同的 WER:

print(wer.compute(predictions=predictions, references=references))

输出:

0.03507271171941831

完美!WER 再次为 3.5%,因为我们的输出与单独使用主模型时完全相同。

投机解码也可以与易用的 🤗 Transformers pipeline API 一起用于推理。下面,我们使用模型和处理器实例化 pipeline,然后用它来转录玩具数据集中的第一个样本。这可以扩展到转录任意长度的音频样本,包括使用批处理:

from transformers import pipeline

pipe = pipeline(
    "automatic-speech-recognition",
    model=model,
    tokenizer=processor.tokenizer,
    feature_extractor=processor.feature_extractor,
    max_new_tokens=128,
    chunk_length_s=15,
    batch_size=4,
    generate_kwargs={"assistant_model": assistant_model},
    torch_dtype=torch_dtype,
    device=device,
)

sample = dataset[0]["audio"]
result = pipe(sample)
print(result["text"])

输出:

 Mr. Quilter is the apostle of the middle classes and we are glad to welcome his gospel.

使用 Whisper 和 Distil-Whisper 运行投机解码的端到端代码片段可以在 Distil-Whisper 模型卡 上找到。它将本 notebook 中涵盖的推理阶段合并为一个单一的代码示例。

多语言语音转录

Distil-Whisper 是英语语音转录的完美助手模型,因为它在短音频和长音频样本上的性能与原始 Whisper 模型的 WER 相差不到 1%,同时速度快 6 倍。然而,官方的 Distil-Whisper 检查点仅支持英语,这意味着它们无法用于多语言语音转录。

要使用推测解码进行多语言语音转录,可以使用 官方多语言 Whisper 检查点之一,或 Whisper 的微调变体。在撰写本文时,Hugging Face Hub 上有超过 5,000 个微调 Whisper 检查点,涵盖 100 多种语言。这些为选择在单一语言上表现优异的助手 Whisper 检查点提供了极好的起点。在本示例中,我们将使用最小的官方多语言检查点 Whisper tiny。欢迎尝试用你的语言微调的不同检查点!

让我们加载新助手模型 Whisper tiny 的权重。由于 Whisper tiny 中的编码器与 large-v2 中的不同,这次我们将使用 AutoModelForSpeechSeq2Seq 类加载编码器和解码器:

assistant_model_id = "openai/whisper-tiny"

assistant_model = AutoModelForSpeechSeq2Seq.from_pretrained(
    assistant_model_id,
    torch_dtype=torch_dtype,
    low_cpu_mem_usage=True,
    use_safetensors=True,
    attn_implementation="sdpa",
)

assistant_model.to(device);

对于我们的基准测试数据集,我们将从 VoxPopuli 数据集的荷兰语("nl")拆分中加载 73 个样本:

dataset = load_dataset("sanchit-gandhi/voxpopuli_dummy", "nl", split="validation")

太好了!我们现在可以像之前一样重新运行基线 Whisper large-v2 模型的基准测试。我们唯一改变的是将语言和任务参数传递给生成函数,以确保我们执行的是语音转录(而非语音翻译)。推测解码与语音转录和翻译任务完全兼容。只需根据需要设置任务参数,如下所示:

all_time = 0
predictions = []
references = []

for sample in tqdm(dataset):
    audio = sample["audio"]
    inputs = processor(audio["array"], sampling_rate=audio["sampling_rate"], return_tensors="pt")
    inputs = inputs.to(device=device, dtype=torch.float16)
    
    output, gen_time = generate_with_time(model, inputs, language="nl", task="transcribe")
    all_time += gen_time
    predictions.append(processor.batch_decode(output, skip_special_tokens=True, normalize=True)[0])
    references.append(processor.tokenizer._normalize(sample["normalized_text"]))

wer_result = wer.compute(predictions=predictions, references=references)

print("Time:", all_time)
print("WER:", wer_result)

输出:

100%|██████████| 73/73 [02:05<00:00,  1.72s/it]
Time: 116.50992178916931
WER: 0.127190136275146

没错!我们的基线时间为 117 秒,WER 为 12.8%。让我们使用推测解码重新运行生成过程:

all_time = 0
predictions = []
references = []

for sample in tqdm(dataset):
    audio = sample["audio"]
    inputs = processor(audio["array"], sampling_rate=audio["sampling_rate"], return_tensors="pt")
    inputs = inputs.to(device=device, dtype=torch.float16)

    output, gen_time = assisted_generate_with_time(model, inputs, language="nl", task="transcribe")
    all_time += gen_time
    predictions.append(processor.batch_decode(output, skip_special_tokens=True, normalize=True)[0])
    references.append(processor.tokenizer._normalize(sample["normalized_text"]))

wer_result = wer.compute(predictions=predictions, references=references)

print("Time:", all_time)
print("WER:", wer_result)

输出:

100%|██████████| 73/73 [01:08<00:00,  1.06it/s]
Time: 62.10229682922363
WER: 0.127190136275146

再次,我们实现了 12.8% 的 WER,但这次推理时间仅为 62 秒,加速了 1.9 倍。鉴于加载助手模型的开销很低,并且数学上保证输出完全相同,推测解码为现有 Whisper 流水线提供了完美的即插即用替代方案。

高效推测解码的策略

在最后一节中,我们介绍两种策略,以确保使用推测解码实现最快的推理时间。

助手模型

我们的目标是选择一个助手模型,它至少比主模型快 3 倍并且至少正确转录 70-80% 的预测 token,通常是示例中“较容易”的 token。如果你有想要转录的特定语言,一个有效的策略是训练两个不同大小的 Whisper 模型,并将一个用作另一个的助手:

  • 首先,微调 Whisper large-v3 作为你的主模型
  • 其次,在同一数据集上蒸馏 Whisper large-v3 作为快速助手模型

微调和蒸馏可以改善主模型和助手模型在你所选语言上的 WER 性能,同时最大化 token 分布的对齐。Whisper 微调的完整指南可以在这里找到,蒸馏的指南在这里。

批量大小

值得注意的是,投机解码在批大小为 1 时速度提升最大。对于批量投机解码,整个批次中的所有候选 token 都必须与验证 token 匹配,token 才会被接受。如果批次中某个位置的 token 不一致,该位置之后的所有候选 token 都会被丢弃。因此,投机解码更有利于较小的批大小。在实践中,我们发现投机解码在批大小不超过 4 时能提供加速。当批大小超过 4 时,投机解码的推理速度比单独使用主模型更慢。完整结果请参阅 Distil-Whisper 论文的 D.3 节。

结论

在这篇博客文章中,我们介绍了投机解码的推理策略,并将其应用于用于语音转录的 Whisper 模型。我们展示了如何实现 2 倍加速,同时在数学上确保与单独使用原始模型时相同的输出。鉴于使用额外辅助模型的开销很低,并且能保证相同的转录结果,我们鼓励你尝试将投机解码作为现有 Whisper 流水线的直接替代方案。

致谢

博客文章由 Sanchit Gandhi 撰写。非常感谢 Patrick von Platen 和 Pedro Cuenca 提出的建设性意见,以及 Joao Gante 在 🤗 Transformers 中实现的辅助生成功能。

来源:Hugging Face:Blog(RSS) · huggingface.co