ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

PaddleSpeech ErnieLinear 标点恢复模型:从 API 结构到训练、推理与源码剖析

PaddleSpeech ErnieLinear 标点恢复模型:从 API 结构到训练、推理与源码剖析 人工智能语音音频【免费下载链接】PaddleSpeechEasy-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.项目地址https://gitcode.com/gh_mirrors/pa/PaddleSpeech点击查看免费下载导读本文围绕 PaddleSpeech 文本模块中的paddlespeech.text.models.ernie_linear.ernie_linear模块展开系统讲解其核心类ErnieLinear的 API 设计、前向计算逻辑、配套数据集与训练/评估组件并结合examples/iwslt2012/punc0的完整实战流程数据准备、训练、测试、标点恢复推理以及仓库中相关源码帮助读者掌握在 PaddleSpeech 中基于 ERNIE 预训练模型训练标点恢复Punctuation Restoration模型并落地推理的完整技术方案。读完本文你将能够理解该模块的输入输出契约、配置文件各参数的含义与调参要点并能够复现或扩展一个基于 ERNIE 线性分类头的标点恢复系统。模块定位与 API 文档说明ErnieLinear模块位于仓库 paddlespeech/text/models/ernie_linear/ 目录下对应 API 文档为 docs/source/api/paddlespeech.text.models.ernie_linear.ernie_linear.rst。该 RST 文档使用 Sphinx 的automodule指令自动生成 API 参考.. automodule:: paddlespeech.text.models.ernie_linear.ernie_linear :members: :undoc-members: :show-inheritance:其中:members:表示展示模块内所有公开成员:undoc-members:表示即使没有 docstring 的成员也会被列出:show-inheritance:会显示类的继承关系。也就是说这份 API 文档的内容完全由 paddlespeech/text/models/ernie_linear/ernie_linear.py 中的公开类与函数自动生成真正值得研究的是该模块的源码实现。从包结构看paddlespeech/text/models/ernie_linear/init.py 依次导出了三个子模块的内容dataset.py提供两个标点数据集类PuncDataset、PuncDatasetFromErnieTokenizerernie_linear.py提供核心模型类ErnieLinearernie_linear_updater.py提供训练器ErnieLinearUpdater与评估器ErnieLinearEvaluator。因此这份 API 文档虽然只有短短几行指令但其背后对应着一整套预训练 ERNIE 线性分类头的标点恢复实现这也是本文展开的主线。ErnieLinear 模型结构与前向计算构造参数与两种初始化路径ErnieLinear继承自paddle.nn.Layer构造函数签名如下见 ernie_linear.pyclass ErnieLinear(nn.Layer): def __init__(self, num_classesNone, pretrained_tokenernie-1.0, cfg_pathNone, ckpt_pathNone, **kwargs):核心参数含义num_classes分类标签数。对标点恢复任务而言标签集合通常是无标点/空格、逗号、句号、问号等因此配置中常取 4。注意该参数仅在未提供cfg_path与ckpt_path的路径下被校验要求必须是大于 0 的整数。pretrained_tokenPaddleNLP 预训练模型名默认ernie-1.0。在examples/iwslt2012/punc0/conf/的多个配置中该值被替换为ernie-3.0-base-zh、ernie-3.0-mini-zh、ernie-tiny等不同规模的 ERNIE 3.0 系列模型详见下文配置小节。cfg_path/ckpt_path本地配置目录与 checkpoint 路径。二者必须同时提供或同时为空用于加载本地微调后的模型。构造函数内部存在两条初始化分支对应两种使用场景加载本地模型cfg_path与ckpt_path均不为None先将路径展开为绝对路径并校验文件存在随后调用ErnieForTokenClassification.from_pretrained(os.path.dirname(cfg_path))即从cfg_path所在目录加载本地模型配置与权重。从头微调预训练模型默认分支校验num_classes为合法整数后调用ErnieForTokenClassification.from_pretrained(pretrained_token, num_labelsnum_classes, **kwargs)在 ERNIE 预训练模型之上叠加 token 分类头输出维度为num_labels。无论走哪条分支模型都会记录self.num_classes self.ernie.num_labels并实例化一个paddle.nn.Softmax()层用于输出概率分布。这里的关键点在于ErnieForTokenClassification来自 PaddleNLPpaddlenlp.transformers它把 ERNIE 编码器与一个线性 token 分类头封装在一起PaddleSpeech 的ErnieLinear在其之上再做一次 reshape 与 softmax从而统一了训练需要 logits与推理需要概率两种场景的输出形态。forward 输入输出与 logits/softmax 双输出设计forward方法见 ernie_linear.py的输入是标准的 BERT/ERNIE 风格四件套input_idstoken 化后的输入 id 序列token_type_ids句子片段 id用于区分句对的两段也可用于区分前一 token 与当前 token的边界信息position_ids位置编码 id可选attention_mask注意力掩码可选。前向流程为y self.ernie(input_ids, token_type_idstoken_type_ids, attention_maskattention_mask, position_idsposition_ids) y paddle.reshape(y, shape[-1, self.num_classes]) logits self.softmax(y) return y, logits即先由 ERNIE 编码器输出每个 token 的原始分类 logits形状为[batch, seq_len, num_classes]然后paddle.reshape将其展平为[-1, num_classes]-1自动推断为batch * seq_len最后对最后一维做 softmax 得到概率。返回值是(y, logits)二元组y是未归一化的 logits用于训练时的交叉熵损失logits是 softmax 概率用于推理时的argmax取预测标签。从使用方代码可以印证这一契约训练器 ernie_linear_updater.py 中loss self.criterion(y, label)使用第一个返回值计算CrossEntropyLoss推理脚本 punc_restore.py 中preds paddle.argmax(logits, axis-1)使用第二个返回值取标签。配套数据集PuncDataset 与 PuncDatasetFromErnieTokenizer标点恢复需要去标点文本 → 标点标签的序列标注数据。模块提供了两个数据集类见 dataset.py二者都遵循paddle.io.Dataset接口__getitem__返回(input_data, label)。PuncDataset词级词典映射PuncDataset(train_path, vocab_path, punc_path, seq_len100)面向词级别建模通过load_vocab读取词表vocab_path与标点表punc_path并为词表预留UNK、END两个特殊 token为标点表预留空格 表示无标点标签通常对应 id 0读取训练文本后按空白切分成 token 序列逐个 token 判断其下一个token 是否为标点从而构造当前词 → 标点标签的训练样本词表中不存在的词回退为UNK样本按seq_len默认 100截断并 reshape 成[-1, seq_len]的定长块长度不足的部分直接丢弃。从源码看self.in_len len(input_data) // self.seq_len决定了数据集的有效样本数也就是说该数据集要求文本总长度至少达到一个seq_len块数据量不足时数据集可能为空这是使用该数据集时需要注意的前提。PuncDatasetFromErnieTokenizertoken 级切分PuncDatasetFromErnieTokenizer(train_path, punc_path, pretrained_tokenernie-1.0, seq_len100)面向ERNIE tokenizer 的 token 级别建模这也是examples/iwslt2012/punc0/conf/default.yaml中dataset_type: Ernie实际使用的数据集通过ErnieTokenizer.from_pretrained(pretrained_token)加载与模型配套的 tokenizer记录pad_token_id预处理时对每个词调用self.tokenizer(word)并取input_ids[1:-1]即去掉 tokenizer 自动添加的[CLS]与[SEP]两个特殊 token只保留该词真正的子词subwordid 序列因为一个词可能被切成多个子词标签对齐规则为某个词的最后一个子词承接该词后面真实的标点标签其余子词标签为无标点空格 id即for i in range(len(x) - 1): label.append(self.punc2id[ ])最后一个子词再根据下一个 token 是否标点打标签最终同样按seq_len截断并 reshape返回 numpy 数组由paddle.io.DataLoader自动转 tensor。这个子词级别的标签对齐是 ERNIE 标点恢复训练中最容易出错的地方PuncDatasetFromErnieTokenizer.preprocess的实现保证了每个输入子词都有且仅有一个标签且标签与 token 序列严格等长。训练与评估组件ErnieLinearUpdater / ErnieLinearEvaluatorernie_linear_updater.py 复用了 PaddleSpeech T2S 训练框架paddlespeech.t2s.training的StandardUpdater与StandardEvaluator为标点任务定制了训练/评估逻辑。ErnieLinearUpdater的关键成员与行为构造参数依次为model、criterion、scheduler、optimizer、dataloader、output_dir其中 criterion 为paddle.nn.CrossEntropyLoss由train.py的DefinedLoss注册表映射每个进程rank独立写worker_{rank}.log日志文件便于多卡训练时按进程排查update_core中执行完整的一步训练labelreshape 成[-1]、模型前向得到(y, logit)、argmax得到预测、交叉熵损失反向传播、optimizer.step()与scheduler.step()推进每个 batch 都会用sklearn.metrics.f1_score(..., averagemacro)计算宏平均 F1并通过report(train/loss, ...)、report(train/F1_score, ...)上报指标供VisualDL等扩展记录。ErnieLinearEvaluator结构与前者对称evaluate_core在验证集上计算 loss 与宏平均 F1 并上报eval/loss、eval/F1_score不进行梯度更新。这种训练器评估器Trainer 扩展的组织方式使得整个训练流程可以通过配置驱动在 train.py 中DefinedClassifier模型、DefinedLoss损失、DefinedDataset数据集三张注册表把 YAML 配置中的model_type、loss_type、dataset_type字符串映射到具体实现从而一份配置即可完整描述数据 → 模型 → 优化器 → 训练器全链路。实战基于 IWSLT2012-Zh 的标点恢复全流程仓库在 examples/iwslt2012/punc0/ 提供了可直接运行的标点恢复示例run.sh通过--stage/--stop-stage控制四个阶段见 run.sh./run.sh --stage 0 --stop-stage 0 # 数据准备 ./run.sh --stage 1 --stop-stage 1 # 模型训练 ./run.sh --stage 2 --stop-stage 2 # 测试输出 classification_report ./run.sh --stage 3 --stop-stage 3 # 标点恢复单条文本推理run.sh顶部的关键变量gpus0,1 conf_pathconf/default.yaml train_output_pathexp/default ckpt_namesnapshot_iter_12840.pdz text今天的天气真不错啊你下午有空吗我想约你一起去吃饭text是一段无标点中文文本阶段 3 会为它补上逗号、句号、问号是最直观的演示入口。Stage 0数据准备./local/data.sh负责下载并切分 IWSLT2012-Zh 语料生成data/iwslt2012_zh/下的train.txt、dev.txt、test.txt以及标点词表punc_vocab。其中train.txt等文件是词与标点交替的序列文本格式供上述两个数据集类直接消费。Stage 1模型训练./local/train.sh conf/default.yaml exp/default实际调用python3 ${BIN_DIR}/train.py \ --configconf/default.yaml \ --output-direxp/default \ --ngpu1train.pypaddlespeech/text/exps/ernie_linear/train.py的主要流程为根据ngpu与是否编译 CUDA 决定设备ngpu0时使用 CPU否则使用 GPUworld_size 1时初始化并行环境seed_everything(config.seed)固定随机种子分别用DefinedDataset[config[dataset_type]]构建训练集与验证集Ernie类型即PuncDatasetFromErnieTokenizer并封装DataLoader训练集shuffleTruemodel DefinedClassifier[config[model_type]](**config[model])即ErnieLinear(num_classes..., pretrained_token...)优化器为Adamweight_decay由L2Decay注入学习率调度为ExponentialDecay(learning_rate, gamma)组装ErnieLinearUpdater与Trainer注册ErnieLinearEvaluator每 1 个 epoch 触发、VisualDL每 1 个 iteration 触发与Snapshot每 1 个 epoch 触发扩展最终trainer.run()。训练产物checkpoint、日志、VisualDL 数据均输出到--output-dir指定的目录。Stage 2测试评估./local/test.sh conf/default.yaml exp/default snapshot_iter_12840.pdz调用 test.pypython3 ${BIN_DIR}/test.py \ --configconf/default.yaml \ --checkpointexp/default/checkpoints/snapshot_iter_12840.pdz \ --print_evalTrue \ --ngpu1test.py加载state_dict[main_params]恢复模型权重注意 checkpoint 是嵌套字典main_params是模型参数所在键在test_path上逐 batch 前向并argmax取预测最后输出 sklearn 的classification_report当--print_eval为 True 时还会额外输出按标点类别统计的 Precision / Recall / F1 表格。evaluation函数按labels[1, 2, 3]即除无标点外的逗号、句号、问号三类分别统计并汇总OVERALLmacro 平均。Stage 3单条文本标点恢复./local/punc_restore.sh调用 punc_restore.pypython3 ${BIN_DIR}/punc_restore.py \ --configconf/default.yaml \ --checkpointexp/default/checkpoints/snapshot_iter_12840.pdz \ --text今天的天气真不错啊你下午有空吗我想约你一起去吃饭其推理流程值得细读_clean_text将输入转为小写并去掉除中文、英文字母、数字以外的字符以及标点表内已有标点得到纯文本preprocess用ErnieTokenizer(list(clean_text), return_lengthTrue, is_split_into_wordsTrue)将每个汉字/字符作为独立 token 切分记录input_ids、seg_idstoken_type_ids与真实长度seq_len模型前向得到 logitsargmax取预测标签再通过tokenizer.convert_ids_to_tokens还原每个子词对应的 token拼接时对每个 token若标签非 0无标点则在其后插入对应标点最终打印Punctuation Restoration Result:结果。以run.sh内置的演示文本为例预期输出大致为今天的天气真不错啊你下午有空吗我想约你一起去吃饭。——注意实际效果取决于所用 checkpoint 与模型规模。配置文件深度解读与调参要点examples/iwslt2012/punc0/conf/下提供了多份配置default.yamlernie-1.0、ernie-3.0-base.yaml、ernie-3.0-medium.yaml、ernie-3.0-mini.yaml、ernie-3.0-nano-zh.yaml、ernie-tiny.yaml结构完全一致。以 default.yaml 为例逐段解读数据段DATA SETTINGdataset_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: 100dataset_typeErnie对应PuncDatasetFromErnieTokenizer子词级Punc对应PuncDataset词级data_params.pretrained_token必须与模型段的pretrained_token一致否则 tokenizer 与模型词表不匹配会导致 id 错位seq_len每个训练块的长度默认 100。更大的seq_len提供更长上下文但会增大显存占用与单样本计算量数据量小于seq_len时会因整除截断而产生空数据集需要留意punc_path标点词表路径内容形如每行一个标点符号空格 被自动追加为无标点标签id 0。模型段MODEL SETTINGmodel_type: ErnieLinear model: pretrained_token: ernie-1.0 num_classes: 4model_type对应注册表中的ErnieLinearnum_classes: 4对应标签集合无标点、逗号、句号、问号4 类。若标点词表扩充如加入顿号、冒号、分号num_classes需同步调整为len(punc_vocab) 1。优化器段OPTIMIZER / SCHEDULER SETTINGoptimizer_params: weight_decay: 1.0e-6 scheduler_params: learning_rate: 1.0e-5 gamma: 0.9999weight_decayAdam 的 L2 权重衰减系数默认 1e-6learning_rate初始学习率 1e-5。ERNIE 预训练模型微调通常使用远小于从头训练的较小学习率避免破坏预训练权重gammaExponentialDecay的衰减因子源码注释明确要求必须在 (0.0, 1.0) 之间且越接近 1.0 越好配合max_epoch: 20可在整个训练期内保持平稳衰减。训练与其他段TRAINING / OTHER SETTINGmax_epoch: 20 num_snapshots: 10 seed: 42max_epoch训练总轮数num_snapshotsSnapshot扩展最多保留的 checkpoint 数超出后自动淘汰旧快照seed随机种子经seed_everything同时固定paddle、random与numpy保证多进程训练可复现源码注释强调multiprocess training 的必需项。模型规模选择与参考效果除了default.yaml中的ernie-1.0仓库还提供了 ERNIE 3.0 全系列配置ernie-3.0-base-zh、ernie-3.0-medium-zh、ernie-3.0-mini-zh、ernie-3.0-micro-zh、ernie-3.0-nano-zh与ernie-tiny。这些配置与 examples/iwslt2012/punc0/README.md 中给出的各模型测试结果一一对应可据此按精度优先/速度优先选择Ernie 1.0base 规模整体宏平均 F1 约 0.633问号类 F1 最高约 0.841适合追求更高分类精度的场景ERNIE 3.0 系列base / medium / mini / micro / nano从 base 到 nano 模型体积递减base 整体 F1 约 0.680 为系列最高nano 约 0.575模型越小Precision 与 Recall 之间的差距越明显Recall 明显下降ernie-tiny整体 F1 约 0.618是轻量部署移动端/嵌入式的折中选择。需要说明的是以上数值来自仓库 README 中记录的 IWSLT2012-Zh 测试集结果属于特定数据集、特定训练配置下的参考值实际复现时因数据版本、训练时长与超参不同会存在波动。测试结果表格中三类标点COMMA / PERIOD / QUESTION的 Precision、Recall、F1 以及OVERALLmacro 平均与test.py中evaluation函数的统计口径一致。与 PaddleSpeech 其他模块的协作关系从源码结构看ErnieLinear并非孤立模块它与 PaddleSpeech 的多个子系统存在协作训练框架复用ErnieLinearUpdater/ErnieLinearEvaluator继承自paddlespeech.t2s.training.updaters.standard_updater.StandardUpdater与paddlespeech.t2s.training.extensions.evaluator.StandardEvaluator并借助paddlespeech.t2s.training.trainer.Trainer以及Snapshot、VisualDL等扩展完成训练闭环体现了文本任务复用 TTS 训练基建的设计推理落地punc_restore.py的标点恢复能力被demos/punctuation_restoration/run.sh与 CLI 标点命令包装为可直接调用的服务也常见于 ASR 后处理流水线——ASR 输出的无标点转写文本经ErnieLinear补标点后可显著提升下游 TTS 合成与文本阅读体验模型注册机制train.py/test.py/punc_restore.py均通过DefinedClassifier {ErnieLinear: ErnieLinear}这类注册表按配置字符串实例化模型新增模型只需在注册表中添加映射即可扩展成本低。总结paddlespeech.text.models.ernie_linear.ernie_linear是 PaddleSpeech 文本标点恢复能力的核心 API。本文从其自动生成的 RST 文档出发完整梳理了ErnieLinear的两种初始化路径与(logits, softmax)双输出前向契约、两类数据集的标签对齐差异、训练/评估组件的工作方式并以examples/iwslt2012/punc0为实例走通了数据准备 → 训练 → 测试 → 推理的四个阶段最后逐段解读了 YAML 配置中的调参要点与不同 ERNIE 规模的选型参考。读者既可以据此复现 IWSLT2012-Zh 标点恢复基线也可以将ErnieLinear直接嵌入自己的 ASR 后处理或文本预处理管线或替换为其他 PaddleNLP 预训练模型进行扩展实验。赞分享人工智能语音音频【免费下载链接】PaddleSpeechEasy-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.项目地址https://gitcode.com/gh_mirrors/pa/PaddleSpeech点击查看免费下载相关推荐PaddleSpeech 标点恢复Punctuation Restoration实战指南从命令行推理到 ErnieLinear 模型原理PaddleSpeech 标点恢复Punctuation Restoration实战指南从命令行推理到 ErnieLinear 模型原理 标点恢复Pun人工智能语音音频NLP媒体生成LunaTranslator完整教程3种模式把日文游戏变成母语从配置到排错一次讲清LunaTranslator完整教程3种模式把日文游戏变成母语从配置到排错一次讲清 打开一款日文视觉小说对话框里全是平假名剧情只能靠猜。LunaTr人工智能语音音频NLP媒体生成PaddleSpeech DeepSpeech2 模型模块源码剖析从 CRNN 编码器到训练、解码与推理导出PaddleSpeech DeepSpeech2 模型模块源码剖析从 CRNN 编码器到训练、解码与推理导出 本篇文章以 PaddleSpeech 仓库中 p人工智能语音音频创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表