自动语音识别


文档摘要

自动语音识别 自动语音识别(Automatic Speech Recognition,ASR)把 spoken audio 转换为书面文字,在人类语音与机器可读语言之间架起桥梁。本文件涵盖 GMM-HMM、CTC 损失、RNN-T 转录器、基于注意力的编码器-解码器模型(LAS)、Whisper 以及端到端 ASR——从经典流水线到现代神经架构。 自动语音识别(ASR)的任务是把口语音频转换为书面文字。它是 AI 中最古老的问题之一(20 世纪 50 年代的最早系统只能识别单个数字),也是商业部署最广泛的技术之一(语音助手、转录服务、字幕生成)。 困难之处在于语音的巨大可变性:不同的说话人、口音、语速、背景噪声、麦克风特性,以及把连续声学信号映射到离散词语这一根本性的歧义。

自动语音识别

自动语音识别(Automatic Speech Recognition,ASR)把 spoken audio 转换为书面文字,在人类语音与机器可读语言之间架起桥梁。本文件涵盖 GMM-HMM、CTC 损失、RNN-T 转录器、基于注意力的编码器-解码器模型(LAS)、Whisper 以及端到端 ASR——从经典流水线到现代神经架构。

  • 自动语音识别(ASR)的任务是把口语音频转换为书面文字。它是 AI 中最古老的问题之一(20 世纪 50 年代的最早系统只能识别单个数字),也是商业部署最广泛的技术之一(语音助手、转录服务、字幕生成)。

  • 困难之处在于语音的巨大可变性:不同的说话人、口音、语速、背景噪声、麦克风特性,以及把连续声学信号映射到离散词语这一根本性的歧义。

  • 可以把 ASR 想象成法庭上的速记员。速记员听到连续的声音流,在脑中把它切分成词,利用上下文消除歧义("they're" 还是 "their" 还是 "there"),然后打出结果。ASR 系统做的是同样的事情,只不过分成若干阶段,这些阶段可以显式地分开并独立或联合地优化。

  • 经典 ASR 流水线以一条由若干独立阶段组成的链来处理音频:原始音频被转换为特征(MFCC 或对数梅尔频谱图,来自第 1 个文件),声学模型为每一特征帧与每一语音单元的匹配程度打分,发音模型(词典)把语音单元映射到词,语言模型为词序列的可能性打分,解码器搜索使组合得分最大化的词序列。每个组件都单独训练和调优。

ASR 流水线:从原始音频经过特征提取、声学模型、解码器、语言模型到输出文字

  • 音素是区分语言中词语的最小语音单位。英语大约有 39-44 个音素(确切数目取决于方言和所使用的音素清单)。例如,"bat" 和 "pat" 只在一个音素上不同(/b/ 对 /p/)。大多数 ASR 系统建模的是上下文相关音素,称为三音子(triphones):由左右邻居定义的音素(例如,"b_t" 语境中的 "a" 与 "c_t" 语境中的 "a" 是不同的单元),因为音素的声学实现深受其邻居影响(这被称为协同发音 coarticulation)。

  • 可能的三音子数量极其庞大(40 的三次方 = 64,000),所以决策树聚类把声学上相似的三音子归并为森音(senones,通常 2000-10,000 个类别)。每个森音有自己的声学模型。这种聚类是第 6 章决策树算法的一种形式。

  • GMM-HMM(高斯混合模型 - 隐马尔可夫模型)是 20 世纪 80 年代到 2010 年代初占主导地位的声学建模方法。HMM(来自第 5 章)建模语音的时间结构:每个音素是一个从左到右的、含 3-5 个状态的 HMM,每个状态代表一个子音素段(起始、中间、结束)。状态之间的转移隐式地建模了时长。

  • 在每个 HMM 状态,发射概率(给定该状态时某一特征向量的可能性)由高斯混合模型(GMM)建模:若干多元高斯分布的加权和(来自第 5 章):

p(\mathbf{x} | s) = \sum_{m=1}^{M} w_m \cdot \mathcal{N}(\mathbf{x} ; \boldsymbol{\mu}_m, \boldsymbol{\Sigma}_m)
  • 其中 \mathbf{x} 是特征向量(如 39 维 MFCC),s 是 HMM 状态,M 是混合分量的数目(通常 8-64),w_m 是混合权重,\boldsymbol{\mu}_m\boldsymbol{\Sigma}_m 是每个高斯分量的均值和协方差。协方差矩阵通常是对角的,以提高计算效率(假设各特征维独立,对 MFCC 来说由于 DCT 去相关,这一假设近似成立)。

  • 训练使用 Baum-Welch 算法(EM 的一个特例,来自第 5 章)从带转录的语音数据中迭代估计 GMM 参数和 HMM 转移概率。解码(寻找最可能的状态序列)使用维特比算法(动态规划,来自第 5 章):

\delta_t(j) = \max_{i} \left[ \delta_{t-1}(i) \cdot a_{ij} \right] \cdot b_j(\mathbf{x}_t)
  • 其中 \delta_t(j) 是在时刻 t 结束于状态 j 的最佳路径的概率,a_{ij} 是从状态 i 转移到状态 j 的概率,b_j(\mathbf{x}_t) 是特征 \mathbf{x}_t 在状态 j 下的发射概率。

  • DNN-HMM(Hinton 等,2012)用深度神经网络(DNN,来自第 6 章)取代了 GMM 发射模型,该网络从一窗口的特征帧预测森音后验概率 p(s | \mathbf{x})。HMM 仍然负责时间结构和排序,但神经网络提供了更具区分性的发射得分。这种混合方法相对 GMM 把词错率降低了 20-30%,是 2012-2016 年的主导范式。

  • WFST 解码(加权有限状态转换器,Weighted Finite-State Transducer)是传统 ASR 的标准解码框架。每个组件(HMM 拓扑 H、上下文依赖 C、词典 L、文法/语言模型 G)都表示为一个加权有限状态转换器,并把它们组合成一个单一的搜索图 H \circ C \circ L \circ G。然后维特比搜索在这个组合图中寻找代价最低的路径。WFST 允许知识源的模块化组合和高效的动态规划搜索。其数学框架来自有限自动机理论(与第 5 章的状态机相关)。

  • 端到端 ASR 消除了各独立组件(发音模型、音素清单、WFST 解码器),训练单个神经网络直接从音频特征映射到字符或词片。关键挑战是对齐问题:输入(每秒数百个特征帧)和输出(每秒几个字符)长度差异很大,而且训练时它们之间的对齐是未知的。

  • 连接时序分类(Connectionist Temporal Classification,CTC)(Graves 等,2006)通过引入一个特殊的空白(blank)token 来解决对齐问题,允许网络输出任意由字符和空白组成的序列,只要把连续重复字符合并并去掉空白后能得到正确的转录。例如,转录 "cat" 可以由输出序列 "--cc-aa-t--"(其中 "-" 是空白)产生。

  • 形式上,CTC 定义了一个多对一映射 \mathcal{B},从所有长度为 T 的输出序列(字母表加空白上的序列)映射到标签序列。标签序列 \mathbf{y} 的概率是所有坍缩为它的对齐之和:

P(\mathbf{y} | \mathbf{x}) = \sum_{\boldsymbol{\pi} \in \mathcal{B}^{-1}(\mathbf{y})} \prod_{t=1}^{T} p(\pi_t | \mathbf{x})

CTC 对齐:穿过空白和字符 token 的许多可能路径都坍缩为相同的输出文字

  • 朴素地计算这个和需要枚举指数级多的对齐,但 CTC 前向-后向算法用动态规划在 O(T \cdot |\mathbf{y}|) 时间内高效完成,类似于第 5 章的 HMM 前向-后向算法。

  • CTC 做了一个条件独立性假设:给定输入,每个时间步的输出独立于所有其他输出。这意味着 CTC 无法建模输出之间的依赖(例如,它无法学到 "q" 几乎总是后接 "u")。必须使用外部语言模型来处理这类依赖。

  • CTC 解码选项:

    • 贪心解码:在每个时间步取最可能的 token,然后合并。快但次优。
    • 束搜索:在每一步保留前 k 个部分假设,合并坍缩为相同前缀的假设。可以融入语言模型得分。
    • 前缀束搜索:一种改进的束搜索,正确处理 CTC 的空白合并,确保假设在合并后进行比较。
  • RNN 转录器(RNN-Transducer,RNN-T)(Graves,2012)通过添加一个显式的预测网络(一个类似语言模型的 RNN)扩展了 CTC,该网络把每个输出条件化于之前的输出之上,去除了条件独立性假设。RNN-T 有三个组件:

    • 编码器:处理音频特征,产生隐藏表示 \mathbf{h}_t^\text{enc}(通常是若干层 LSTM 或 Conformer)。
    • 预测网络:一个自回归 RNN,从之前已发出的标签产生隐藏表示 \mathbf{h}_u^\text{pred}
    • 联合网络:在每个(时间,标签)位置组合编码器和预测网络的输出,产生下一个 token(包括空白)上的分布:
p(y | t, u) = \text{softmax}(W \cdot \text{tanh}(W_\text{enc} \mathbf{h}_t^\text{enc} + W_\text{pred} \mathbf{h}_u^\text{pred} + b))
  • RNN-T 在每个时间步可以发出零个或多个标签(通过在推进到下一个时间步之前发出非空白 token,或通过发出空白来推进但不输出)。训练在二维(时间,标签)格上使用前向-后向算法,复杂度为 O(T \cdot U),其中 U 是输出长度。RNN-T 是设备端流式 ASR 的主导架构(用于 Google 的 Pixel 手机及类似产品),因为它天然支持流式处理:编码器从左到右处理音频,预测网络增量地生成输出。

  • 听、注意和拼(Listen, Attend and Spell,LAS)(Chan 等,2016)是一个基于注意力的编码器-解码器模型(第 6 章的序列到序列架构)。它有三个组件:

    • 听者(编码器):一个金字塔形双向 LSTM,处理整个输入序列并通过因子 8 下采样(在每一层通过拼接相邻隐藏状态对),产生更短的编码器隐藏状态序列。
    • 注意力:在每个解码器步骤,计算对所有编码器状态的注意力权重以形成上下文向量(与第 7 章的注意力机制相同)。
    • 拼写者(解码器):一个自回归 LSTM,以一个字符一个字符地生成输出转录,条件化于上下文向量和之前生成的字符。
  • LAS 取得了不错的结果,但要求在解码前整个话语都可用(因为注意力会关注所有编码器状态),这使它不适合流式应用。它对非常长的话语也表现不佳,因为对长序列的注意力会变得分散。

  • Conformer(Gulati 等,2020)把卷积捕捉局部模式的能力与自注意力建模全局依赖的能力结合起来。每个 Conformer 块有四个模块,呈三明治结构:

    1. 前馈模块(半步):一个带残差连接的前馈网络,残差权重减半。
    2. 多头自注意力模块:标准的 Transformer 自注意力(来自第 7 章),配以相对位置编码。
    3. 卷积模块:一个逐点卷积、一个门控线性单元(GLU)、一个一维深度卷积、批归一化、一个 Swish 激活,以及另一个逐点卷积。深度卷积捕捉局部上下文(类似特征序列上的 n-gram)。
    4. 前馈模块(半步):与模块 1 相同。
  • 输出为:\mathbf{y} = \text{LayerNorm}(\mathbf{x} + \frac{1}{2}\text{FFN}_1 + \text{MHSA} + \text{Conv} + \frac{1}{2}\text{FFN}_2)。这种马卡龙式的结构(FFN-Attention-Conv-FFN)配以半步残差,是经验上发现优于其他排列的。Conformer 已成为 CTC 和 RNN-T 系统的默认编码器,性能优于纯 Transformer 和纯 LSTM 编码器。

Conformer 块:前馈、自注意力、卷积、前馈模块的三明治结构

  • Whisper(Radford 等,2023)是 OpenAI 的大规模基于注意力的 ASR 模型。它使用标准的编码器-解码器 Transformer 架构(来自第 7 章),在 680,000 小时的弱监督数据(从互联网上抓取的、音频配对近似转录)上训练。关键设计选择:

    • 输入:80 通道对数梅尔频谱图(来自第 1 个文件),25 ms 窗、10 ms 步长,归一化为零均值单位方差。
    • 编码器:标准 Transformer 编码器,配正弦位置嵌入和前激活层归一化。
    • 解码器:Transformer 解码器,使用字节级 BPE 分词器(来自第 7 章)自回归地生成 token。
    • 多任务:单个模型处理转录、翻译、语言识别和时间戳预测,由解码器提示中的特殊任务 token 条件化。
    • 训练数据的规模(而非架构上的新颖性)是 Whisper 在跨领域、跨口音、跨语言上强泛化的主要驱动力。
  • wav2vec 2.0(Baevski 等,2020)是一个用于语音表示的自监督预训练框架。核心思想是从大量无标注音频中学习语音表示,然后用少量标注数据微调。这遵循与 BERT(来自第 7 章)相同的自监督范式,但适配于连续的音频信号。

  • wav2vec 2.0 架构有三部分:

    • 特征编码器:一个多层一维 CNN,处理原始波形样本,以 20 ms 的帧率产生潜在表示 \mathbf{z}_t(16 kHz 下每 320 个样本一个向量)。
    • 量化模块:使用乘积量化(把向量分成若干组,每组独立量化,从 G 个各有 V 项的码本中选择)把潜在表示离散化为有限码本。这为对比学习目标产生目标 \mathbf{q}_t
    • 上下文网络:一个 Transformer 编码器,接收(部分被掩码的)潜在表示并产生上下文化的表示 \mathbf{c}_t

wav2vec 2.0 架构:CNN 特征编码器、掩码、Transformer 上下文网络,以及与量化目标的对比学习

  • 在预训练期间,潜在表示的随机片段被掩码(替换为一个学习到的掩码嵌入),模型必须从一组干扰项(从同一话语的其他位置采样的负样本)中识别出被掩码位置的真实量化表示。对比损失为:
\mathcal{L} = -\log \frac{\exp(\text{sim}(\mathbf{c}_t, \mathbf{q}_t) / \kappa)}{\sum_{\tilde{\mathbf{q}} \in Q_t} \exp(\text{sim}(\mathbf{c}_t, \tilde{\mathbf{q}}) / \kappa)}
  • 其中 \text{sim} 是余弦相似度,\kappa 是温度参数,Q_t 包括真实量化目标加上干扰项。一个额外的多样性损失鼓励均匀使用所有码本项。这个损失本质上是 InfoNCE 对比损失,与视觉自监督学习中使用的对比目标同属一族。

  • 预训练之后,在顶部添加一个线性投影和 CTC 头,并在标注数据上微调。wav2vec 2.0 仅用 10 分钟标注数据(预训练使用 53,000 小时无标注音频)就取得了接近最优的结果,证明了自监督学习在低资源语音识别中的威力。

  • HuBERT(Hsu 等,2021)是另一种自监督方法,它用掩码预测目标(预测被掩码帧的离散聚类分配)取代了对比目标。目标由一个离线聚类步骤产生(第一次迭代对 MFCC 做 k-means,后续迭代对 HuBERT 特征做 k-means)。HuBERT 相对 wav2vec 2.0 简化了训练流水线(不需要量化模块或对比采样),并取得了相当或更好的结果。

  • Fast Conformer(Rekesh 等,2023,NVIDIA NeMo)用下采样注意力机制取代了标准 Conformer 中的平方复杂度自注意力:输入序列在计算注意力之前被压缩(通常通过步幅卷积下采样 8 倍),然后再扩展回去。这将注意力代价从 O(T^2) 降到 O(T^2/64),同时保留全局上下文,使得可以在非常长的话语(长达数分钟)上训练而不出现内存问题。Fast Conformer 是 NVIDIA NeMo 工具包的默认编码器,构成了其生产级模型的骨干。

  • Parakeet(NVIDIA,2024)是一系列基于 Fast Conformer 编码器配 CTC 和 RNN-T 解码器的高精度英语 ASR 模型,在 64,000 小时英语语音上训练。Parakeet 模型(0.6B 和 1.1B 参数)在发布时在标准基准上取得了最低的词错率,在大多数英语测试集上超越了 Whisper large-v3。关键要素是高效的 Fast Conformer 架构、激进的数据增强(SpecAugment、速度扰动、噪声混合)以及大规模监督训练数据——这证明了对已知组件的精心工程化仍然能推动技术前沿。

  • Canary(NVIDIA,2024)把 NeMo 框架扩展到多语言和多任务 ASR。它使用 Fast Conformer 编码器配基于注意力的解码器(而非 CTC 或 RNN-T),在单个模型中处理跨多种语言的转录加翻译(类似 Whisper 的多任务设计,但采用更高效的 Fast Conformer 骨干)。Canary 模型支持英语、德语、西班牙语和法语,准确率有竞争力。

  • Moonshine(Useful Sensors,2024)是一系列专门为设备端和边缘部署优化的 ASR 模型。其编码器使用混合架构,用一个小型 CNN 加几层 Transformer 替换最初的 Transformer/Conformer 层,大幅减小了模型尺寸(基础模型不到 30M 参数)。Moonshine 面向在 CPU 和低功耗设备上的实时流式处理,在这些场景下 Whisper 会太大太慢,它以一定精度换取 5-10 倍更低的延迟和内存占用。

  • Distil-Whisper(Gandhi 等,2023)把知识蒸馏(第 6 章)应用于 Whisper,将其压缩成更小更快的模型。学生模型只使用 2 个解码器层(而 Whisper 有 32 个),同时保留完整编码器,并被训练来匹配 Whisper 的输出分布。Distil-Whisper 在词错率上与教师模型相差不到 1%,同时快 6 倍,使其适用于完整 Whisper 模型太慢的实时应用。

  • 通用语音模型(Universal Speech Model,USM)(Zhang 等,2023,Google)把自监督预训练扩展到 300 多种语言的 1200 万小时无标注音频,随后进行监督微调。USM 证明了 wav2vec 2.0 / 自监督范式可以扩展到真正的大规模数据 regime,在标注数据非常有限的低资源语言上取得强劲表现。

  • 大规模多语言语音(Massively Multilingual Speech,MMS)(Pratap 等,2023,Meta)把 wav2vec 2.0 预训练扩展到 1,100 多种语言,使用宗教录音和其他多语言音频来源。MMS 覆盖的语言远多于以往任何 ASR 系统,首次为许多资源匮乏的语言实现了语音识别。

  • 现代 ASR 的格局正在收敛到几个主导模式:(1) 用于流式的 Conformer 家族编码器配 CTC 或 RNN-T;(2) 用于离线/多任务的编码器-解码器 Transformer;(3) 用于低资源场景的自监督预训练;(4) 规模——更多数据和更大的模型持续提升准确率。它们之间的选择取决于部署约束:延迟预算、可用算力、语言数量,以及应用是流式还是批处理。

  • 语言模型融合通过融入超出声学模型所捕捉范围的语言学知识来改进 ASR。基本思想是在解码时把声学模型得分 p(\mathbf{x} | \mathbf{y})(音频与转录的匹配程度)与语言模型得分 p(\mathbf{y})(转录作为一个句子的可能性)结合起来。

  • 浅融合在束搜索时组合得分:

\hat{\mathbf{y}} = \arg\max_\mathbf{y} \left[ \log p_\text{AM}(\mathbf{y} | \mathbf{x}) + \lambda \log p_\text{LM}(\mathbf{y}) \right]
  • 其中 \lambda 是可调权重,p_\text{LM} 是外部语言模型(通常是第 7 章的 n-gram 或神经语言模型)。这种方法简单有效,但要求语言模型在与 ASR 模型相同的 token 词表上运行。

  • 深融合(Gulcehre 等,2015)把语言模型集成到解码器网络内部:语言模型的隐藏状态与解码器隐藏状态拼接,并通过一个门控机制后再进行输出投影。整个系统(包括预训练的语言模型)被联合微调。这允许更深度的集成,但训练更复杂。

  • 冷融合(Sriram 等,2018)类似深融合,但它是从头训练 ASR 解码器并集成语言模型,而不是微调一个预训练的解码器。这迫使声学模型学习与语言模型互补的信息,而不是复制语言模型已经知道的内容。

  • 重打分(N-best rescoring)是一种两遍方法:首先用束搜索生成 N 个候选转录,然后用一个更强大的语言模型(例如大型 Transformer 语言模型)对它们重新排序。这易于实现,并允许使用对于第一遍解码来说太慢的非常大的语言模型。

  • 内部语言模型估计(Internal Language Model Estimation,ILME)解决一个微妙的问题:端到端模型会从训练转录中隐式地学习一个内部语言模型,这在浅融合时可能与外部语言模型冲突(本质上是重复计算语言先验)。ILME 估计内部语言模型并在融合时减去其得分:

\hat{\mathbf{y}} = \arg\max_\mathbf{y} \left[ \log p_\text{E2E}(\mathbf{y} | \mathbf{x}) - \beta \log p_\text{ILM}(\mathbf{y}) + \lambda \log p_\text{LM}(\mathbf{y}) \right]
  • 流式与离线 ASR 是一个根本性的架构选择。离线(或批处理)ASR 在产生任何输出之前处理整个话语。流式 ASR 随着音频到达增量地产生输出,延迟有界。

  • 流式对实时应用至关重要:实时字幕、语音助手(用户期望在自己说完之前就得到响应)、电话通话转录。挑战在于,一些未来的上下文对识别是有帮助的(知道下一个词是 "York" 可以消除 "New" 的歧义),但流式系统不能等待任意长的未来上下文。

  • 单向编码器(从左到右的 LSTM、因果卷积、因果 Transformer)天然支持流式,因为每个输出只依赖于过去和现在的输入。双向编码器(会查看未来上下文)不直接支持流式。

  • 分块注意力(也称块注意力或段注意力)把输入分成固定长度的块,并只在每个块内(可选地包括之前的几个块)施加自注意力。这将延迟限制为块大小加处理时间,同时仍允许在每个块内有一些局部的双向上下文。权衡是准确率随块大小减小而下降。

  • 前瞻允许流式编码器在为当前帧产生输出之前偷看少量未来帧(例如 300-900 ms)。这通过在单向计算中添加少量右上下文来实现。前瞻窗口增加了延迟,但显著提升准确率。

  • 流式 ASR 中的延迟有几个组成部分:

    • 算法延迟:从音频到达到模型能够处理它的延迟(由块大小、前瞻和特征提取决定)。
    • 计算延迟:运行模型前向传播的时间。
    • 端点检测延迟:检测用户已经说完的延迟。
    • 首 token 延迟:第一个词出现的速度。最终确认延迟:最终输出被确认的速度(流式系统通常产生临时输出,随着更多音频到达而被更正)。
  • ASR 的评估指标

  • 词错率(Word Error Rate,WER)是主要指标。它通过使用编辑距离(把一个转换成另一个所需的最少替换、插入和删除次数)把假设(系统输出)与参考(真实转录)对齐来计算,然后:

\text{WER} = \frac{S + D + I}{N}
  • 其中 S 是替换、D 是删除、I 是插入,N 是参考中的总词数。如果插入很多,WER 可能超过 100%。对于清晰的朗读语音,5% 的 WER 被认为大致是人类水平;对话或嘈杂语音要困难得多(10-20%+)。

  • 字符错率(Character Error Rate,CER)是同一公式应用于字符级而非词级。对于没有清晰词边界的语言(中文、日文),CER 更有信息量,也用于评估近似错误的接近程度("cat" 对 "bat" 是 100% WER 但只有 33% CER)。

  • 词信息丢失(Word Information Lost,WIL)和词信息保留(Word Information Preserved,WIP)是信息论意义上的替代指标,比 WER 更精确地考虑参考与假设之间的相关性,但报告得较少。

  • 实时因子(Real-Time Factor,RTF)衡量计算效率:处理时间与音频时长的比值。RTF < 1 表示系统比实时快;RTF > 1 表示无法跟上实时音频。流式系统必须维持 RTF < 1。

  • 数据增强对鲁棒的 ASR 至关重要。常用技术:

    • 速度扰动:以 0.9 倍和 1.1 倍速度对音频重采样(改变音高和时长)。
    • SpecAugment(Park 等,2019):在频谱图中掩码随机频带和时间步。这是音频版的 dropout,是 ASR 最有效的正则化技术之一。它不需要额外数据。
    • 噪声增强:以各种信噪比把干净语音与录制的噪声混合。
    • 房间脉冲响应仿真:把干净语音与模拟的房间声学做卷积,以仿真混响环境。
  • ASR 的分词决定模型的输出词表。选项包括:

    • 字符:简单,词表小(英语约 30 个),但输出序列长,且没有隐式的语言建模。
    • 词片 / BPE(来自第 7 章):在词表大小和序列长度之间取得平衡的子词单元。现代系统的标准(Whisper 使用字节级 BPE,约 50,000 个 token)。
    • :词表大(50,000+),输出序列短,但无法处理词表外的词。
    • 音素:有语言学依据、紧凑,但需要发音词典。
  • ASR 的演进可以总结为从重度工程的模块化系统(GMM-HMM + WFST 解码,1990s-2010s)到混合系统(DNN-HMM,2012-2016)到把越来越多的流水线吸收进单个神经网络的端到端系统(CTC、RNN-T、LAS,2016-2020)再到利用海量无标注或弱标注数据的大规模预训练模型(wav2vec 2.0、Whisper,2020 至今)的进步过程。每一次转变都简化了工程并提高了准确率,遵循了机器学习中从数据学习表示而非手工设计表示的更广泛趋势(与第 6 章中图像特征被 CNN 取代、第 7 章中 NLP 特征被 Transformer 取代是同一个故事)。

编程练习(使用 CoLab 或 notebook)

  1. 在 JAX 中从零实现 CTC 损失。创建一个由短序列 logits 和一个目标标签组成的小例子,计算 CTC 前向算法得到总概率,并计算负对数似然损失。
import jax import jax.numpy as jnp import matplotlib.pyplot as plt def ctc_forward(log_probs, targets): """ CTC 前向算法(对数域以保证数值稳定性)。 log_probs: (T, V) 词表上的对数概率(索引 0 = 空白) targets: (U,) 目标标签索引(不含空白) 返回:CTC 下目标序列的对数概率。 """ T, V = log_probs.shape U = len(targets) # 构造带空白的扩展标签序列:[blank, y1, blank, y2, ..., yU, blank] S = 2 * U + 1 labels = jnp.zeros(S, dtype=jnp.int32) # 全部为空白 for i in range(U): labels = labels.at[2 * i + 1].set(targets[i]) # 初始化 alpha(对数域) NEG_INF = -1e30 alpha = jnp.full((T, S), NEG_INF) alpha = alpha.at[0, 0].set(log_probs[0, labels[0]]) # 从空白开始 alpha = alpha.at[0, 1].set(log_probs[0, labels[1]]) # 或从第一个标签开始 # 填充前向变量 for t in range(1, T): for s in range(S): # 同一状态 a = alpha[t - 1, s] # 来自前一状态 if s > 0: a = jnp.logaddexp(a, alpha[t - 1, s - 1]) # 跳过空白(当当前和前一个的前一个标签不同时) if s > 1 and labels[s] != 0 and labels[s] != labels[s - 2]: a = jnp.logaddexp(a, alpha[t - 1, s - 2]) alpha = alpha.at[t, s].set(a + log_probs[t, labels[s]]) # 总对数概率:最后时间步最后两个状态之和 log_prob = jnp.logaddexp(alpha[T - 1, S - 1], alpha[T - 1, S - 2]) return log_prob, alpha # --- 小例子 --- T = 12 # 输入长度(时间步) V = 5 # 词表大小(0=空白, 1='c', 2='a', 3='t', 4='x') targets = jnp.array([1, 2, 3]) # "c", "a", "t" # 创建随机 logits 并转换为对数概率 key = jax.random.PRNGKey(42) logits = jax.random.normal(key, (T, V)) log_probs = jax.nn.log_softmax(logits, axis=-1) log_prob, alpha = ctc_forward(log_probs, targets) ctc_loss = -log_prob print(f"Target sequence: {targets.tolist()} ('c', 'a', 't')") print(f"Input length T={T}, Vocab size V={V}") print(f"CTC log-probability: {log_prob:.4f}") print(f"CTC loss (neg log-prob): {ctc_loss:.4f}") # 可视化前向变量(alpha)格 fig, ax = plt.subplots(figsize=(12, 5)) # 从对数转换为线性以便可视化 alpha_linear = jnp.exp(alpha - jnp.max(alpha)) # 归一化以便观察 im = ax.imshow(alpha_linear.T, aspect='auto', origin='lower', cmap='viridis') ax.set_xlabel('Time step (t)') ax.set_ylabel('Extended label index (s)') label_names = ['_', 'c', '_', 'a', '_', 't', '_'] # _ = 空白 ax.set_yticks(range(len(label_names))) ax.set_yticklabels(label_names) ax.set_title(f'CTC Forward Variable (alpha lattice) | Loss = {ctc_loss:.2f}') plt.colorbar(im, ax=ax, label='Normalised probability') plt.tight_layout(); plt.show()
  1. 在 JAX 中构建一个简单的编码器-解码器基于注意力的 ASR 模型(一个最小的类 LAS 架构)。使用一维卷积编码器和带点积注意力的单层解码器。在合成数据上运行并可视化注意力权重。
import jax import jax.numpy as jnp import matplotlib.pyplot as plt # --- 最小的基于注意力的编码器-解码器 ASR --- def init_params(key, input_dim, hidden_dim, vocab_size): """初始化一个微小的类 LAS 模型的参数。""" keys = jax.random.split(key, 8) scale = 0.1 params = { # 编码器:简单的线性投影(模拟卷积输出) 'enc_w': jax.random.normal(keys[0], (input_dim, hidden_dim)) * scale, 'enc_b': jnp.zeros(hidden_dim), # 注意力:query、key、value 投影 'attn_q': jax.random.normal(keys[1], (hidden_dim, hidden_dim)) * scale, 'attn_k': jax.random.normal(keys[2], (hidden_dim, hidden_dim)) * scale, 'attn_v': jax.random.normal(keys[3], (hidden_dim, hidden_dim)) * scale, # 解码器 RNN(为说明用简单的 Elman RNN) 'dec_wh': jax.random.normal(keys[4], (hidden_dim, hidden_dim)) * scale, 'dec_wx': jax.random.normal(keys[5], (vocab_size, hidden_dim)) * scale, 'dec_wc': jax.random.normal(keys[6], (hidden_dim, hidden_dim)) * scale, 'dec_b': jnp.zeros(hidden_dim), # 输出投影 'out_w': jax.random.normal(keys[7], (hidden_dim, vocab_size)) * scale, 'out_b': jnp.zeros(vocab_size), } return params def encode(params, x): """编码器:线性投影(占位,代表 conv/LSTM 堆栈)。""" return jnp.tanh(x @ params['enc_w'] + params['enc_b']) def attend(params, query, enc_out): """对编码器输出做点积注意力。""" q = query @ params['attn_q'] # (hidden,) k = enc_out @ params['attn_k'] # (T_enc, hidden) v = enc_out @ params['attn_v'] # (T_enc, hidden) d_k = q.shape[-1] scores = (k @ q) / jnp.sqrt(d_k) # (T_enc,) weights = jax.nn.softmax(scores) # (T_enc,) context = weights @ v # (hidden,) return context, weights def decode_step(params, h_prev, y_prev_onehot, enc_out): """单个解码器步:RNN + 注意力。""" # 嵌入前一个 token y_emb = y_prev_onehot @ params['dec_wx'] # (hidden,) # 对编码器做注意力 context, attn_w = attend(params, h_prev, enc_out) # RNN 更新 h = jnp.tanh(h_prev @ params['dec_wh'] + y_emb + context @ params['dec_wc'] + params['dec_b']) # 输出 logits logits = h @ params['out_w'] + params['out_b'] return h, logits, attn_w # --- 设置 --- key = jax.random.PRNGKey(0) input_dim = 40 # 例如 40 个梅尔带 hidden_dim = 64 vocab_size = 10 # 演示用的小词表 T_enc = 30 # 编码器时间步 T_dec = 8 # 解码器步数 params = init_params(key, input_dim, hidden_dim, vocab_size) # 合成输入:随机的类梅尔特征 key, subkey = jax.random.split(key) x = jax.random.normal(subkey, (T_enc, input_dim)) # 编码 enc_out = encode(params, x) # 解码(教师强制,使用随机目标) key, subkey = jax.random.split(key) targets = jax.random.randint(subkey, (T_dec,), 0, vocab_size) h = jnp.zeros(hidden_dim) all_logits = [] all_attn = [] for t in range(T_dec): y_prev = jax.nn.one_hot(targets[t] if t > 0 else 0, vocab_size) h, logits, attn_w = decode_step(params, h, y_prev, enc_out) all_logits.append(logits) all_attn.append(attn_w) all_attn = jnp.stack(all_attn) # (T_dec, T_enc) all_logits = jnp.stack(all_logits) # (T_dec, vocab_size) # --- 可视化注意力权重 --- fig, axes = plt.subplots(1, 2, figsize=(14, 5)) im = axes[0].imshow(all_attn, aspect='auto', cmap='Blues', origin='lower') axes[0].set_xlabel('Encoder time step') axes[0].set_ylabel('Decoder step') axes[0].set_title('Attention Weights (decoder -> encoder)') plt.colorbar(im, ax=axes[0]) # 展示每个解码器步的预测 token 分布 im2 = axes[1].imshow(jax.nn.softmax(all_logits, axis=-1), aspect='auto', cmap='Oranges', origin='lower') axes[1].set_xlabel('Vocabulary index') axes[1].set_ylabel('Decoder step') axes[1].set_title('Output Token Probabilities') plt.colorbar(im2, ax=axes[1]) plt.suptitle('Minimal Attention-based ASR Model (untrained)') plt.tight_layout(); plt.show()
  1. 使用动态规划(编辑距离)从零计算词错率(WER),并针对一个参考评估多个假设。可视化编辑距离矩阵。
import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np def compute_wer(reference, hypothesis): """ 使用动态规划(词级莱文斯坦距离)计算 WER。 返回 WER、替换数、删除数、插入数以及 DP 矩阵。 """ ref_words = reference.split() hyp_words = hypothesis.split() N = len(ref_words) M = len(hyp_words) # DP 矩阵:d[i][j] = ref[:i] 与 hyp[:j] 之间的编辑距离 d = np.zeros((N + 1, M + 1), dtype=np.int32) # 回溯矩阵以统计 S、D、I ops = np.zeros((N + 1, M + 1, 3), dtype=np.int32) # [sub, del, ins] for i in range(N + 1): d[i][0] = i # 全部为删除 for j in range(M + 1): d[0][j] = j # 全部为插入 for i in range(1, N + 1): for j in range(1, M + 1): if ref_words[i - 1] == hyp_words[j - 1]: sub_cost = d[i - 1][j - 1] # 匹配,无编辑 else: sub_cost = d[i - 1][j - 1] + 1 # 替换 del_cost = d[i - 1][j] + 1 # 删除 ins_cost = d[i][j - 1] + 1 # 插入 d[i][j] = min(sub_cost, del_cost, ins_cost) # 回溯以统计操作 i, j = N, M S, D, I = 0, 0, 0 while i > 0 or j > 0: if i > 0 and j > 0 and d[i][j] == d[i-1][j-1] and ref_words[i-1] == hyp_words[j-1]: i -= 1; j -= 1 # 正确 elif i > 0 and j > 0 and d[i][j] == d[i-1][j-1] + 1: S += 1; i -= 1; j -= 1 # 替换 elif i > 0 and d[i][j] == d[i-1][j] + 1: D += 1; i -= 1 # 删除 elif j > 0 and d[i][j] == d[i][j-1] + 1: I += 1; j -= 1 # 插入 else: break wer = (S + D + I) / N if N > 0 else 0.0 return wer, S, D, I, d # --- 测试用例 --- reference = "the cat sat on the mat" hypotheses = [ "the cat sat on the mat", # 完美 "the cat sit on the mat", # 1 个替换 "the cat on the mat", # 1 个删除 "the big cat sat on the mat", # 1 个插入 "a dog sat in a rug", # 多个错误 ] print(f"Reference: '{reference}'\n") print(f"{'Hypothesis':<40s} {'WER':>6s} {'S':>3s} {'D':>3s} {'I':>3s}") print("-" * 60) results = [] for hyp in hypotheses: wer, S, D, I, dp = compute_wer(reference, hyp) results.append((hyp, wer, S, D, I, dp)) print(f"'{hyp}':<40s} {wer:>6.1%} {S:>3d} {D:>3d} {I:>3d}") # 可视化最差情况的 DP 矩阵 worst = results[-1] hyp_words = worst[0].split() ref_words = reference.split() dp_matrix = worst[5] fig, axes = plt.subplots(1, 2, figsize=(14, 5)) # DP 矩阵 im = axes[0].imshow(dp_matrix, cmap='YlOrRd', origin='upper') axes[0].set_xticks(range(len(hyp_words) + 1)) axes[0].set_xticklabels([''] + hyp_words, rotation=45, ha='right', fontsize=9) axes[0].set_yticks(range(len(ref_words) + 1)) axes[0].set_yticklabels([''] + ref_words, fontsize=9) axes[0].set_xlabel('Hypothesis words') axes[0].set_ylabel('Reference words') axes[0].set_title(f'Edit Distance Matrix\nWER = {worst[1]:.1%}') for i in range(dp_matrix.shape[0]): for j in range(dp_matrix.shape[1]): axes[0].text(j, i, str(dp_matrix[i, j]), ha='center', va='center', fontsize=8) plt.colorbar(im, ax=axes[0]) # WER 对比柱状图 names = [f'Hyp {i+1}' for i in range(len(results))] wers = [r[1] * 100 for r in results] colors = ['#27ae60' if w == 0 else '#f39c12' if w < 30 else '#e74c3c' for w in wers] axes[1].barh(names, wers, color=colors) axes[1].set_xlabel('WER (%)') axes[1].set_title('Word Error Rate Comparison') for i, (w, r) in enumerate(zip(wers, results)): axes[1].text(w + 1, i, f'{w:.0f}% (S={r[2]}, D={r[3]}, I={r[4]})', va='center', fontsize=9) axes[1].set_xlim(0, max(wers) * 1.4) plt.tight_layout(); plt.show()
  1. 在对数梅尔频谱图上实现 SpecAugment(频率掩码和时间掩码),并可视化原始版本与增强版本的对比。从合成信号生成频谱图。
import jax import jax.numpy as jnp import matplotlib.pyplot as plt # --- 生成合成的对数梅尔频谱图 --- key = jax.random.PRNGKey(42) fs = 16000 duration = 2.0 t = jnp.arange(0, duration, 1.0 / fs) # 模拟语音:带谐波的啁啾信号 f0 = 120.0 x = sum(jnp.sin(2 * jnp.pi * f0 * k * t * (1 + 0.1 * t)) / k for k in range(1, 10)) key, subkey = jax.random.split(key) x = x + 0.05 * jax.random.normal(subkey, t.shape) # 计算对数梅尔频谱图(简化版) frame_len = 400 # 25 ms hop_len = 160 # 10 ms n_fft = 512 n_mels = 80 n_frames = (len(x) - frame_len) // hop_len + 1 hamming = 0.54 - 0.46 * jnp.cos(2 * jnp.pi * jnp.arange(frame_len) / (frame_len - 1)) frames = jnp.stack([x[i * hop_len : i * hop_len + frame_len] for i in range(n_frames)]) windowed = frames * hamming spectra = jnp.abs(jnp.fft.rfft(windowed, n=n_fft)) ** 2 # 简单的梅尔滤波器组 def hz_to_mel(f): return 2595 * jnp.log10(1 + f / 700) def mel_to_hz(m): return 700 * (10 ** (m / 2595) - 1) mel_points = jnp.linspace(hz_to_mel(0), hz_to_mel(fs / 2), n_mels + 2) hz_pts = mel_to_hz(mel_points) bins = jnp.floor((n_fft + 1) * hz_pts / fs).astype(jnp.int32) n_freqs = n_fft // 2 + 1 fb = jnp.zeros((n_mels, n_freqs)) for m in range(n_mels): lo, mid, hi = int(bins[m]), int(bins[m+1]), int(bins[m+2]) for k in range(lo, mid): if mid != lo: fb = fb.at[m, k].set((k - lo) / (mid - lo)) for k in range(mid, hi): if hi != mid: fb = fb.at[m, k].set((hi - k) / (hi - mid)) log_mel = jnp.log(spectra @ fb.T + 1e-10) # --- SpecAugment --- def spec_augment(spec, key, n_freq_masks=2, freq_mask_width=15, n_time_masks=2, time_mask_width=25): """应用 SpecAugment:频率和时间掩码。""" augmented = spec.copy() T, F = spec.shape # 频率掩码 for _ in range(n_freq_masks): key, k1, k2 = jax.random.split(key, 3) f_width = jax.random.randint(k1, (), 1, freq_mask_width + 1) f_start = jax.random.randint(k2, (), 0, max(1, F - freq_mask_width)) mask = (jnp.arange(F) >= f_start) & (jnp.arange(F) < f_start + f_width) augmented = jnp.where(mask[None, :], 0.0, augmented) # 时间掩码 for _ in range(n_time_masks): key, k1, k2 = jax.random.split(key, 3) t_width = jax.random.randint(k1, (), 1, time_mask_width + 1) t_start = jax.random.randint(k2, (), 0, max(1, T - time_mask_width)) mask = (jnp.arange(T) >= t_start) & (jnp.arange(T) < t_start + t_width) augmented = jnp.where(mask[:, None], 0.0, augmented) return augmented key, subkey = jax.random.split(key) log_mel_aug = spec_augment(log_mel, subkey) # --- 可视化 --- fig, axes = plt.subplots(2, 1, figsize=(14, 8)) im0 = axes[0].imshow(log_mel.T, aspect='auto', origin='lower', cmap='inferno', extent=[0, duration, 0, n_mels]) axes[0].set_title('Original Log-Mel Spectrogram') axes[0].set_xlabel('Time (s)'); axes[0].set_ylabel('Mel Band') plt.colorbar(im0, ax=axes[0], label='Log Energy') im1 = axes[1].imshow(log_mel_aug.T, aspect='auto', origin='lower', cmap='inferno', extent=[0, duration, 0, n_mels]) axes[1].set_title('After SpecAugment (frequency + time masking)') axes[1].set_xlabel('Time (s)'); axes[1].set_ylabel('Mel Band') plt.colorbar(im1, ax=axes[1], label='Log Energy') plt.tight_layout(); plt.show()

作者与出处
原作者: HenryNdubuaku
来源:HenryNdubuaku
许可证:Apache-2.0
整理: 灏天文库整理
由灏天文库结构化整理,提供目录导航、全文检索与在线阅读,便于系统化学习
发布者: 作者: HenryNdubuaku 转发
评论区 (0)
U