- 人工智能
- 语音
- 音频
【免费下载链接】PaddleSpeech
Easy-to-use Speech Toolkit including Self-Supervised Learning model, SOTA/Streaming ASR with punctuation, Streaming TTS with text frontend, Speaker Verification System, End-to-End Speech Translation and Keyword Spotting. Won NAACL2022 Best Demo Award.
导读
本文以 PaddleSpeech 文本任务中的标点恢复(Punctuation Restoration)核心模块paddlespeech.text.models.ernie_linear为主线,完整解析其三大组成部分——基于 ERNIE 预训练模型的 Token 分类器ErnieLinear、两类数据处理集PuncDataset/PuncDatasetFromErnieTokenizer,以及训练/评估器ErnieLinearUpdater/ErnieLinearEvaluator。通过阅读本文,你将掌握该模块的模型结构、数据切分与对齐原理、训练与评估指标计算方式,并能够结合 examples/iwslt2012/punc0 示例跑通"数据准备 → 训练 → 测试 → 标点恢复"的完整流程。
模块定位:ernie_linear在 PaddleSpeech 文本任务中的角色
在 PaddleSpeech 的 API 文档体系中,docs/source/api/paddlespeech.text.models.ernie_linear.rst是该模块的自动文档索引页,它通过automodule指令挂载了模块内所有公开成员,并列出三个子模块:
paddlespeech.text.models.ernie_linear.dataset——数据加载与序列化;paddlespeech.text.models.ernie_linear.ernie_linear——模型本体;paddlespeech.text.models.ernie_linear.ernie_linear_updater——训练与评估逻辑。
从源码结构看,该模块承担的是对标点符号进行 Token 级序列标注的任务:给定一段没有标点的中文(或英文)文本,模型在每个字/词单元后预测是否应该插入逗号、句号、问号等标点,这与语音识别(ASR)结果的后处理、TTS 前端文本规整等场景密切相关。它的实际消费方包括:
- CLI 工具入口 paddlespeech/cli/text/infer.py 中的
TextExecutor(任务punc); - TTS 韵律预测模块 paddlespeech/t2s/frontend/rhy_prediction/rhy_predictor.py(作为句子切分与停顿预测的前置手段);
- 服务端配置 paddlespeech/server/conf/application.yaml 中的文本服务。
模型实现:ErnieLinear——在 ERNIE 之上叠加线性分类头
构造逻辑与两条初始化路径
模型文件 paddlespeech/text/models/ernie_linear/ernie_linear.py 定义了ErnieLinear(nn.Layer)。构造函数__init__(num_classes=None, pretrained_token='ernie-1.0', cfg_path=None, ckpt_path=None, **kwargs)支持两种初始化方式:
- 本地微调模型路径加载:当同时传入
cfg_path与ckpt_path时,代码会对两者做os.path.abspath(os.path.expanduser(...))规范化并断言文件存在,随后通过ErnieForTokenClassification.from_pretrained(os.path.dirname(cfg_path))从配置文件所在目录加载已微调好的 ERNIE Token 分类模型。这正是 CLI 推理时复用训练产物(config + checkpoint)的方式。 - 从预训练权重新建:否则要求
num_classes为正整数,调用ErnieForTokenClassification.from_pretrained(pretrained_token, num_labels=num_classes, **kwargs),默认预训练模型为ernie-1.0,并通过self.num_classes = self.ernie.num_labels从加载后的模型反向取得真实类别数。
其中ErnieForTokenClassification来自 PaddleNLP(paddlenlp.transformers),即 Paddle 生态的 ERNIE 预训练模型 + Token 分类输出头的组合。
前向计算:logits 与 softmax 双输出
forward(input_ids, token_type_ids=None, position_ids=None, attention_mask=None)的流程为:
- 将输入送入 ERNIE 编码器得到每个 Token 的分类 logits
y; y = paddle.reshape(y, shape=[-1, self.num_classes])把[batch, seq_len, num_classes]展平为[batch*seq_len, num_classes],方便后续与展平后的标签做逐 Token 交叉熵;- 经过
nn.Softmax()得到概率分布logits,与原始 logitsy一起返回。
注意:forward返回元组(y, logits),前者用于损失计算,后者用于argmax预测——这在训练、测试、推理脚本中是一致的调用约定(如y, logit = self.model(input))。
损失与推理细节的印证
- 训练时使用
nn.CrossEntropyLoss(见训练脚本中DefinedLoss = {"ce": nn.CrossEntropyLoss}); - 推理时对
logits执行paddle.argmax(logits, axis=1)得到每个 Token 的标点类别索引,类别 0 表示"无标点",1/2/3 分别对应逗号、句号、问号(与数据集标点词表顺序一致)。
数据集实现:文本的"字-标点"序列化与对齐
数据模块 paddlespeech/text/models/ernie_linear/dataset.py 导出两个paddle.io.Dataset子类:PuncDataset与PuncDatasetFromErnieTokenizer。
原始训练文本的格式约定
两者都从train_path读取 UTF-8 文本文件,并按空白切分为词序列:
tmp_seqs = open(train_path, encoding='utf-8').readlines() self.txt_seqs = [i for seq in tmp_seqs for i in seq.split()]即"每个词后面跟一个标点符号(或空格占位)"的交替排列。例如训练语料中形如我 今天 去 公园 , 你 呢 ?,preprocess按位置奇偶解析出(词, 标点)对:若当前 Token 是标点则跳过;下一个 Token 在标点词表中则记为对应标签,否则记为空格(无标点)标签。
PuncDataset:基于自定义词表
PuncDataset(train_path, vocab_path, punc_path, seq_len=100):
- 通过
load_vocab构建word2id(词表从vocab_path读取,并追加<UNK>、<END>两个特殊词,占 id 0/1)与punc2id(标点词表从punc_path读取,追加空格" "占 id 0); - 每个词经
word2id.get(token, word2id["<UNK>"])映射为 id,标签为标点 id 或空格 id; - 最后按
seq_len截断并reshape(-1, seq_len),__len__返回样本条数(in_len = len(input_data) // seq_len)。
PuncDatasetFromErnieTokenizer:基于 BPE/子词切分
PuncDatasetFromErnieTokenizer(train_path, punc_path, pretrained_token='ernie-1.0', seq_len=100)是dataset_type: Ernie(默认配置)实际使用的版本:
- 使用
ErnieTokenizer.from_pretrained(pretrained_token)对每个词调用tokenizer(word)并取input_ids[1:-1](去掉[CLS]/[SEP]),即一个词可能被切分成多个子词 Token; - 对同一词的多个子词 Token,仅最后一个子词后可能挂标点,其余子词一律打"无标点"标签(
for i in range(len(x)-1): label.append(self.punc2id[" "])),从而保证输入 Token 数与标签数严格一一对应(assert len(input_data) != len(label)用于校验); - 输出为
np.array形态的[N, seq_len]整数张量,paddingID = tokenizer.pad_token_id被保留供填充使用。
这个对齐细节是理解模型训练的关键:标签的数量必须与 ERNIE 分词后 Token 的数量完全一致,否则无法做逐 Token 的交叉熵损失。
训练与评估器:ErnieLinearUpdater与ErnieLinearEvaluator
文件 paddlespeech/text/models/ernie_linear/ernie_linear_updater.py 基于 PaddleSpeech TTS 训练框架(paddlespeech.t2s.training)的StandardUpdater/StandardEvaluator实现了两个组件。
训练更新器
ErnieLinearUpdater(model, criterion, scheduler, optimizer, dataloader, output_dir)的核心update_core(batch):
- 解包
(input, label),将标签展平为[-1]; y, logit = self.model(input)得到 logits 与 softmax 概率,pred = paddle.argmax(logit, axis=1)得到预测类别;- 用交叉熵
self.criterion(y, label)计算损失,执行optimizer.clear_grad()→loss.backward()→optimizer.step()→scheduler.step(); - 用
sklearn.metrics.f1_score(..., average="macro")计算宏平均 F1,并通过report("train/loss", ...)、report("train/F1_score", ...)上报给训练框架(VisualDL 可视化与日志依赖这些指标); - 日志写入
output_dir/worker_{rank}.log,按分布式 rank 分文件,支持多卡训练。
评估器
ErnieLinearEvaluator(model, criterion, dataloader, output_dir)的evaluate_core(batch)逻辑与训练一致但不反向传播,上报eval/loss与eval/F1_score,用于每个 epoch 结束后的验证。
与训练框架的衔接
训练脚本 paddlespeech/text/exps/ernie_linear/train.py 展示了完整的装配方式:
- 通过
DefinedClassifier = {'ErnieLinear': ErnieLinear}、DefinedDataset = {'Punc': PuncDataset, 'Ernie': PuncDatasetFromErnieTokenizer}字典实现配置驱动的组件选择; - 训练器
Trainer(updater, (config.max_epoch, 'epoch'), output_dir)管理训练循环; - rank 0 进程额外挂载
evaluator(每 1 epoch 触发)、VisualDL(每 1 iteration 触发),所有进程挂载Snapshot(每 1 epoch 触发,最多保留config.num_snapshots份)。
配置体系:以examples/iwslt2012/punc0为例的完整参数解读
示例 examples/iwslt2012/punc0/conf/default.yaml 是与本模块直接配套的完整配置,五个区块如下:
# DATA SETTING dataset_type: Ernie train_path: data/iwslt2012_zh/train.txt dev_path: data/iwslt2012_zh/dev.txt test_path: data/iwslt2012_zh/test.txt batch_size: 64 num_workers: 2 data_params: pretrained_token: ernie-1.0 punc_path: data/iwslt2012_zh/punc_vocab seq_len: 100 # MODEL SETTING model_type: ErnieLinear model: pretrained_token: ernie-1.0 num_classes: 4 # OPTIMIZER SETTING optimizer_params: weight_decay: 1.0e-6 scheduler_params: learning_rate: 1.0e-5 gamma: 0.9999 # TRAINING SETTING max_epoch: 20 num_snapshots: 5 # OTHER SETTING num_snapshots: 10 seed: 42各参数与本模块源码的对应关系:
| 配置项 | 取值示例 | 作用与源码落点 |
|---|---|---|
dataset_type | Ernie/Punc | 选择PuncDatasetFromErnieTokenizer或PuncDataset(见 train.py 的DefinedDataset) |
data_params.pretrained_token | ernie-1.0 | ERNIE 预训练模型名,同时决定 Tokenizer 与模型底座 |
data_params.punc_path | punc_vocab | 标点词表,第 0 行通常为,、。、?等,标签 id 从 1 开始(0 保留给空格) |
data_params.seq_len | 100 | 每个样本的 Token 序列长度,超出的部分被截断丢弃 |
model.num_classes | 4 | 标点类别数(无标点 + 3 种标点),对应ErnieLinear(num_classes=...) |
optimizer_params.weight_decay | 1.0e-6 | 传入Adam(weight_decay=paddle.regularizer.L2Decay(...)) |
scheduler_params | lr=1e-5, gamma=0.9999 | 构造ExponentialDecay学习率调度器,gamma 需在 (0,1) 之间 |
max_epoch/num_snapshots | 20/10 | 训练轮数与快照保留份数,快照默认名为snapshot_iter_*.pdz |
seed | 42 | 经seed_everything统一设置 paddle/random/np 随机种子,多卡训练必须固定 |
此外,示例目录还提供了多种底座变体配置:ernie-3.0-base.yaml、ernie-3.0-medium.yaml、ernie-3.0-mini.yaml、ernie-3.0-nano-zh.yaml、ernie-tiny.yaml(对应ernie-3.0-base-zh、ernie-3.0-medium-zh等预训练 Token),用于在模型容量与推理速度间权衡。
实战流程:从数据到标点恢复
示例 examples/iwslt2012/punc0/run.sh 通过--stage/--stop-stage控制四个阶段,底层脚本分别调用 paddlespeech/text/exps/ernie_linear/ 下的三个入口。
Stage 0:数据准备
./run.sh --stage 0 --stop-stage 0执行./local/data.sh,产出data/iwslt2012_zh/{train,dev,test}.txt("词+标点"交替格式)与punc_vocab(标点词表)。文本规整逻辑可参考 local/preprocess.py。
Stage 1:模型训练
./run.sh --stage 1 --stop-stage 1实际执行 local/train.sh:
python3 ${BIN_DIR}/train.py \ --config=conf/default.yaml \ --output-dir=exp/default \ --ngpu=1训练脚本 train.py 支持--ngpu指定卡数:ngpu=0时paddle.set_device("cpu");ngpu>1时走dist.spawn(train_sp, (args, config), nprocs=args.ngpu)多进程分布式训练(此时model会被DataParallel包裹,且train.py会把配置文件复制一份到输出目录)。所有 checkpoint 位于exp/default/checkpoints/,典型命名为snapshot_iter_12840.pdz。
Stage 2:测试评估
./run.sh --stage 2 --stop-stage 2执行 local/test.sh,调用 test.py:
python3 ${BIN_DIR}/test.py \ --config=conf/default.yaml \ --checkpoint=exp/default/checkpoints/snapshot_iter_12840.pdz测试脚本加载 checkpoint 中的state_dict["main_params"]恢复模型权重,遍历测试集后输出classification_report(每个标点类别的 Precision/Recall/F1);若--print_eval(默认true,支持str2bool解析)为真,还会输出包含O / COMMA / PERIOD / QUESTION与 OVERALL 宏平均的DataFrame评估表。注意测试脚本中labels=[1,2,3]过滤掉类别 0(无标点),只统计真实标点类别的指标。
示例 RESULTS.md 记录了在 IWLST2012-Zh 测试集上的公开结果,例如ernie-1.0底座的 OVERALL F1 约 0.633、ernie-3.0-base-zh底座的 OVERALL F1 约 0.680(不同底座与配置结果有差异,可作复现基准参考)。
Stage 3:标点恢复推理
./run.sh --stage 3 --stop-stage 3执行 local/punc_restore.sh,调用 punc_restore.py,run.sh 中的默认测试文本为:
text=今天的天气真不错啊你下午有空吗我想约你一起去吃饭推理管线依次为:_clean_text清洗(转小写、剔除不在[A-Za-z0-9\u4e00-\u9fa5]及标点词表内的字符)→ErnieTokenizer逐字切分 →model(input_ids, seg_ids)前向 →argmax得到每 Token 的标点类别 → 逐 Token 拼接文本并在类别非 0 处插入对应标点,最终打印Punctuation Restoration Result: 今天的天气真不错啊,你下午有空吗?我想约你一起去吃饭。这类结果。
CLI 与 Python API:开箱即用的标点恢复
除示例脚本外,该模块还被封装为 CLI 命令与 Python API,入口位于 paddlespeech/cli/text/infer.py 的TextExecutor:
paddlespeech text --task punc --input 今天的天气真不错啊你下午有空吗我想约你一起去吃饭相关参数(源自TextExecutor的 argparse 定义):
| 参数 | 默认值 | 说明 |
|---|---|---|
--task | punc | 当前仅支持punc(标点恢复) |
--model | ernie_linear_p7_wudao | 预训练模型 tag,可选列表来自pretrained_models.py中的ernie_linear_p7_wudao、ernie_linear_p3_wudao、ernie_linear_p3_wudao_fast等(见 paddlespeech/resource/pretrained_models.py) |
--lang | zh | zh/en |
--config/--ckpt_path/--punc_vocab | None | 未指定时自动从资源目录下载默认模型、checkpoint 与词表;指定时走本地加载路径 |
--device | paddle.get_device() | 推理设备 |
内部调用链为:_init_from_path(旧模型,ernie_linear_p7/p3_wudao)通过ErnieLinear(cfg_path=..., ckpt_path=...)直接从模型目录加载;_init_from_path_new(新模型,如ernie_linear_p3_wudao_fast)则用ErnieLinear(**config["model"])重建并set_state_dict(state_dict["main_params"])恢复权重,同时根据模型名是否含fast选择ernie-1.0或ernie-3.0-mini-zh作为 Tokenizer。preprocess→infer→postprocess三步与示例脚本中的推理管线完全一致,postprocess中通过if l != 0: text += self._punc_list[l]完成标点插入。
与姊妹模块的横向对比
paddlespeech.text.models目录下还包含同定位的 ernie_crf 模块。两者都做标点恢复,但解码策略不同:ErnieLinear对每个 Token 独立做 softmax + argmax(点式分类),实现简单、推理快;ernie_crf则在序列层面引入 CRF 转移约束,建模标点标签间的相邻依赖。具体选型时,可在保持数据格式兼容的前提下按精度/速度需求切换(相关对比可见 paddlespeech/cli/README_cn.md 中不同模型的命令行用法)。
小结与扩展阅读
paddlespeech.text.models.ernie_linear是 PaddleSpeech 文本任务中一套完整、可复用的标点恢复实现:模型侧以 PaddleNLP 的 ERNIE Token 分类为底座,数据侧提供"词表级"与"Tokenizer 级"两套序列化方案并严格保证 Token-标签对齐,训练侧复用t2s.training框架实现分布式训练、宏平均 F1 评估与快照管理。围绕它,你还可以继续深入:
- 完整示例与复现基准:examples/iwslt2012/punc0/README.md、RESULTS.md;
- 模块导出与包结构:paddlespeech/text/models/ernie_linear/init.py;
- 姐妹模块 CRF 方案:paddlespeech/text/models/ernie_crf;
- 在 TTS 韵律预测中的实际调用:paddlespeech/t2s/frontend/rhy_prediction/rhy_predictor.py。
掌握本模块后,你即可在此基础上复现标点恢复训练、替换底座模型(ERNIE 1.0 ↔ ERNIE 3.0 系列),或将ErnieLinear集成进自己的 ASR 后处理 / TTS 前端管线。
- 人工智能
- 语音
- 音频
【免费下载链接】PaddleSpeech
Easy-to-use Speech Toolkit including Self-Supervised Learning model, SOTA/Streaming ASR with punctuation, Streaming TTS with text frontend, Speaker Verification System, End-to-End Speech Translation and Keyword Spotting. Won NAACL2022 Best Demo Award.
相关推荐
PaddleSpeech 标点恢复(Punctuation Restoration)实战指南:从命令行推理到 ErnieLinear 模型原理
PaddleSpeech 标点恢复(Punctuation Restoration)实战指南:从命令行推理到 ErnieLinear 模型原理 标点恢复(Pun
人工智能语音音频NLP媒体生成PaddleSpeech 标点恢复实战:paddlespeech.text.exps.ernie_linear 模块训练、评测、推理与模型平均全解析
PaddleSpeech 标点恢复实战:paddlespeech.text.exps.ernie_linear 模块训练、评测、推理与模型平均全解析 导读 本文
人工智能语音音频PaddleSpeech 标点恢复(Punctuation Restoration)实战指南:从命令行到源码原理
PaddleSpeech 标点恢复(Punctuation Restoration)实战指南:从命令行到源码原理 标点恢复(Punctuation Restor
人工智能语音音频
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考