首页
/ mode/models 中 DeepSpeech2 语音识别实战:从 LibriSpeech 预处理到 CTC 训练与 WER 评估的完整流水线

mode/models 中 DeepSpeech2 语音识别实战:从 LibriSpeech 预处理到 CTC 训练与 WER 评估的完整流水线

2026-09-06 17:56:34作者:傅爽业Veleda

DeepSpeech2 是 TensorFlow 生态中经典的端到端自动语音识别(ASR)模型,当前仓库在 research/deep_speech 目录下给出了基于 TensorFlow 1.15.3 / 2.3 的完整实现。本篇以 research/deep_speech/README.md 为主线,覆盖数据集下载与预处理、训练/评估命令、全部命令行参数,并结合源码深入讲解模型结构(2 层卷积 + 5 层双向 RNN + 全连接)、CTC 损失的时间步对齐机制、以及贪心解码器如何计算 WER/CER,帮助读者在本地完整复现这一语音识别流水线。

一、模型总览:DeepSpeech2 是什么

research/deep_speech/README.md 的说明,DeepSpeech2 是一个端到端的深度神经网络 ASR 模型,其网络由 2 个卷积层、5 个双向 RNN 层和 1 个全连接层组成,输入特征为从音频中提取的线性谱图(linear spectrogram),损失函数采用连接时序分类(Connectionist Temporal Classification, CTC)。当前实现参考了原作者的 DeepSpeech 代码与 MLPerf 仓库中的参考实现(README 顶部注明该模块为 "No Maintenance Intended",且标注兼容 TensorFlow 1.15.3 与 2.3 两个版本)。

整个目录的文件职责如下:

文件 职责
deep_speech.py 训练/评估主入口,定义全部命令行参数与 model_fn
deep_speech_model.py DeepSpeech2 网络结构定义
data/download.py 下载 LibriSpeech 语料并预处理为 CSV
data/dataset.py 解析 CSV、构建 tf.data.Dataset、实现批量洗牌
data/featurizer.py 谱图与文本标签特征提取
data/vocabulary.txt 词表:a-z'- 共 28 个字符
decoder.py CTC 贪心解码器与 WER/CER 计算
run_deep_speech.sh 一键跑完整 benchmark 的脚本
requirements.txt Python 依赖:nltk>=3.3pandas>=0.23.3soundfile>=0.10.2sox>=1.3.3

二、运行环境与数据准备

2.1 配置 Python 路径与安装依赖

README 要求先将仓库顶层的 /models 目录加入 Python 路径(因为 deep_speech.py 内部 from official.utils.flags import core as flags_core 依赖仓库顶层的 official 包):

export PYTHONPATH="$PYTHONPATH:/path/to/models"

然后安装共享依赖:

pip3 install -r requirements.txt
# 或
pip install -r requirements.txt

依赖清单见 research/deep_speech/requirements.txt,其中 sox 用于 FLAC 转 WAV,soundfile 用于读取音频,nltk 用于计算编辑距离,pandas 用于生成 CSV。

2.2 下载与预处理 LibriSpeech

python data/download.py
# 参数:
#   --data_dir  数据下载与保存目录,默认 /tmp/librispeech_data

使用 --help / -h 可查看全部参数。从 download.py 源码可以看到,脚本内置了 LibriSpeech(OpenSLR 12)全部 7 个分区的下载地址,并额外提供 --train_only(只下 train-clean-100/360 与 train-other-500)、--dev_only(dev-clean + dev-other)、--test_only(test-clean + test-other)三个布尔开关;不带任何开关时默认下载完整数据集。

预处理的实质工作由 convert_audio_and_split_transcript 完成(download.py):

  1. sox.Transformer 把每个分区的 FLAC 音频逐个转成 WAV;
  2. 逐行解析 .trans.txt 转写文件,把转写文本做 NFKD 归一化、ASCII 化、去首尾空白并转小写;
  3. 生成一个 Tab 分隔的三列 CSVwav_filename(wav 绝对路径)、wav_filesize(字节数)、transcript(转写文本)。

这与 README "Dataset" 一节的描述完全一致:训练数据为 train-clean-100 + train-clean-360(约 13 万条样本),验证集为 dev-clean(约 2.7K 行)。CSV 中的 wav_filesize 并非冗余信息——后续训练流水线直接用它作为音频长短的代理指标来做排序,见下文 4.1 节。

2.3 数据加载:DeepSpeechDataset 与 tf.data 流水线

data/dataset.py 负责把 CSV 变成可训练的 tf.data.Dataset

  • AudioConfig(sample_rate, window_ms, stride_ms, normalize) 承载谱图参数,默认采样率 16000 Hz、窗长 20 ms、帧移 10 ms,并对特征做均值/方差归一化;
  • DatasetConfig 校验 CSV 与词表文件存在性,并持有 sortagrad 开关;
  • DeepSpeechDataset 初始化音频/文本两个 featurizer,并固定 num_feature_bins = 161——即 16 kHz 采样率下 20 ms 汉宁窗对应 161 个线性频率 bins(与 featurizer.pycompute_spectrogram_feature 的输出维度一致);
  • input_fndataset.py)通过 tf.data.Dataset.from_generator 逐条产出 {"features", "input_length", "label_length"} 与标签张量,再用 padded_batch 把 batch 内不等长的时序特征填充到最长样本、prefetch(AUTOTUNE) 加速输入管道。

三、模型结构源码剖析

网络定义集中在 deep_speech_model.pyDeepSpeech2.__call__L143-L176)按如下顺序搭建前向计算图:

  1. Conv1:kernel (41, 11)、stride (2, 2)、32 个 filter、relu6 激活、无 bias,输入前做 (20, 5) 对称 padding;
  2. Conv2:kernel (21, 11)、stride (2, 1)、32 个 filter,padding (10, 5)
  3. 卷积输出 reshape 为 [batch, T', feat_size * 32] 送入 RNN;
  4. 5 层双向 RNN:单元类型由 SUPPORTED_RNNS 决定,支持 gru(默认)、lstmrnn;除第一层外,每层 RNN 前都插入 Batch Normalization(L164-L169is_batch_norm = (layer_counter != 0));
  5. FC 层:最后再做一次 Batch Normalization,接 Dense(num_classes, activation="softmax")use_bias 可控(默认 True)。

几个值得注意的实现细节:

  • _conv_bn_layer 中的对称 padding 是为了保证卷积后序列不会比标签短,源码注释明确写了 "This step is required to avoid issues when RNN output sequence is shorter than the label length"(L77-L82)——这正是 CTC 对"输出时间步不少于标签长度"的要求;
  • BatchNorm 采用 momentum=0.997epsilon=1e-5L30-L32),docstring 特别解释了 momentum 偏大时验证精度收敛更慢,可尝试调小到 0.1 以更快看到评估结果;
  • 输出层的 softmax 激活意味着 logits 张量直接就是概率分布,评估端可直接取 argmax

四、训练与评估入口:deep_speech.py 全参数解析

训练命令(README "Run each step individually" 一节):

python deep_speech.py

deep_speech.py 中的 define_deep_speech_flagsL304-L405)注册了 README 提到的四个核心参数及其余全部超参,汇总如下:

参数 默认值 说明
--model_dir /tmp/deep_speech_model/ 训练 checkpoint 保存目录
--export_dir /tmp/deep_speech_saved_model/ SavedModel 导出目录
--train_data_dir 指向 test-clean CSV 的路径 训练集 CSV 文件路径
--eval_data_dir 同上 评估集 CSV 文件路径
--num_gpus GPU 数量,-1 表示使用全部可用 GPU
--batch_size 128 全局 batch size,多卡时必须是 GPU 数整数倍
--train_epochs 10 训练轮数
--epochs_between_evals (公共参数) 每多少 epoch 评估一次
--seed 1 随机种子
--sample_rate 16000 音频采样率
--window_ms 20 谱图帧长(ms)
--stride_ms 10 谱图帧移(ms)
--vocabulary_file data/vocabulary.txt 词表文件路径
--sortagrad True 首个 epoch 按音频长度排序、不洗牌
--rnn_hidden_size 800 每层 RNN 隐状态维度
--rnn_hidden_layers 5 RNN 层数
--rnn_type gru RNN 单元类型:gru/lstm/rnn
--is_bidirectional True RNN 是否双向
--use_bias True 最后一层 FC 是否使用 bias
--learning_rate 5e-4 Adam 初始学习率
--wer_threshold None 达到该 WER 后停止训练;LibriSpeech 上 MLPerf 参考实现的阈值为 0.23

4.1 sortagrad 与批量洗牌

README "Dataset" 节提到"除第一个 epoch 外,训练数据按 batch 洗牌(当 sortagrad 开启时首个 epoch 除外)"。对应源码是 batch_wise_dataset_shuffle:当 epoch_index == 0sortagrad=True 时,样本保持按 wav_filesize(由 _preprocess_data 解析 CSV 时排序得到)升序排列,使同一 mini-batch 内音频长短相近、减少 padding 浪费并加速早期收敛;之后的每个周期则把样本切成若干"桶"(每桶一个 batch),整桶地随机重排——既打乱样本顺序,又保持桶内长短一致,从而与 padded_batch 的填充策略保持协同。训练主循环在 run_deep_speech 中每个训练周期都调用一次该函数。

4.2 CTC 损失与卷积后的时间步对齐

model_fndeep_speech.py)中训练分支的核心是:

ctc_input_length = compute_length_after_conv(
    tf.shape(features)[1], tf.shape(logits)[1], input_length)
loss = tf.reduce_mean(tf.keras.backend.ctc_batch_cost(
    labels, logits, ctc_input_length, label_length))

由于 batch 内样本被 padding 到同一长度,而卷积 stride (2,2)(2,1) 会把时间维度缩短,CTC 需要知道每条样本真实特征在卷积后对应多少时间步compute_length_after_convL43-L68)用比例关系 ctc_input_length = input_length / max_time_steps * ctc_time_steps 精确反推,再连同 label_length 一起传给 ctc_batch_cost。优化器为 Adam(学习率 --learning_rate),训练 op 用 tf.group(minimize_op, update_ops) 把 BatchNorm 的移动统计量更新一并纳入(L167-L172)。

4.3 评估:贪心解码与 WER / CER

evaluate_modeldeep_speech.py)对评估集每条样本取 probabilities(即 softmax 后的概率),交给 DeepSpeechDecoder 做标准 CTC 贪心解码:对每个时间步取 argmax → 用 itertools.groupby 合并连续重复字符 → 剔除 blank 索引(词表 28 个真实字符,blank 默认索引 28,见 decoder.pyvocabulary.txt 中 a-z、'- 共 28 行)。随后用 nltk.metrics.distance.edit_distance 分别计算字错率(WER)(先按空格分词、把每个唯一词映射为单字符再求编辑距离)和字素错率(CER),两者均除以对应参考长度后对全数据集取平均,最终以 {"WER": ..., "CER": ..., "global_step": ...} 返回并写入 benchmark 日志。若 --wer_threshold 达到阈值,训练主循环立即 break(L298-L301)。

多 GPU 方面,入口通过 distribution_utils.get_distribution_strategy 构建 DistributionStrategy,并强制全局 batch size 必须是 GPU 数的整数倍(per_device_batch_sizeL195-L223),因为 Estimator 场景下需要手动除以卡数得到 per-replica batch。

五、一键 Benchmark:run_deep_speech.sh

README 指出 run_deep_speech.sh 以默认参数跑完整流水线:

sh run_deep_speech.sh

脚本按 4 步执行,且 README 特别提醒:benchmark 的训练集包含 train-clean-100、train-clean-360、train-other-500,评估集包含 dev-clean 与 dev-other。对照脚本源码可以看到各步细节:

  1. Step 1python data/download.py 下载并预处理数据集,各分区 CSV 落在 /tmp/librispeech_data/<分区>/LibriSpeech/<分区>.csv
  2. Step 2:用 head -1 保留表头、sed 1d 去头拼接,合成 train_dataset.csv(三个训练分区)与 eval_dataset.csv(两个 dev 分区);
  3. Step 3:用 awk 逐行调用 soxi -D 读取 wav 时长,过滤掉超过 MAX_AUDIO_LEN=27.0 秒的样本,得到 final_train_dataset.csv / final_eval_dataset.csv——超长样本被剔除是为了控制训练时序列长度与 padding 开销;
  4. Step 4:以 nohup 后台启动训练,关键命令行参数为:
nohup python deep_speech.py \
  --train_data_dir=$final_train_file \
  --eval_data_dir=$final_eval_file \
  --num_gpus=-1 \
  --wer_threshold=0.23 \
  --seed=1 >$log_file 2>&1 &

即使用全部 GPU、以 WER ≤ 0.23 作为停止条件(脚本注释标明这是 MLPerf 参考实现的目标),运行日志落盘到 log_<日期> 文件。

六、适用前提与注意事项

  • 环境前提:该模块面向 TensorFlow 1.15.3 / 2.3,README 顶部徽章标注 "No Maintenance Intended",意味着代码以历史参考实现为主,新环境运行前建议按 requirements.txt 固定依赖版本;
  • 磁盘空间:完整下载 LibriSpeech 全部分区(脚本默认行为)体积较大,可按需使用 --train_only / --dev_only
  • 数据格式约定:训练/评估 CSV 必须是 Tab 分隔三列 wav_filename / wav_filesize / transcript 且首行为表头,wav_filesize 会被用作排序键,手工构造数据时须保证该列正确;
  • 词表与解码器一致性--vocabulary_file 默认指向 data/vocabulary.txt(28 个字符),decoder 的 blank_index=28 与之严格对应;更换词表时需同步调整 blank 索引;
  • 超参调优切入点--rnn_type 可在 gru/lstm/rnn 间切换、--sortagrad 关闭后首 epoch 即随机洗牌、BatchNorm momentum 可按源码注释在 0.997 与更小值(如 0.1)之间权衡验证收敛速度,这些都是源码已预留的调参接口。

综上,research/deep_speech 目录给出了从语料下载、谱图特征工程、CTC 训练到贪心解码评估的端到端 ASR 参考实现:download.py 产出三列 CSV,dataset.py 构建 tf.data 管道并实现 sortagrad 批量洗牌,deep_speech_model.py 定义卷积-双向 GRU-FC 主干,deep_speech.py 以 Estimator 组织训练/评估循环,run_deep_speech.sh 则提供一键复现 MLPerf WER=0.23 目标的完整脚本。

登录后查看全文
热门项目推荐
相关项目推荐