简介:这份资源面向自然语言处理初学者与需要为无标点文本做后处理的开发者,提供一套基于PaddleNLP的预测文本添加标点符号源码,可用于语音识别结果整理、OCR文本还原、字幕断句等场景。压缩包共6个文件,以5个Python脚本和1个txt依赖清单为主,整体约7KB,体量轻便,便于快速阅读与二次修改。其中脚本承担模型推理、日志记录与线性层封装等职责,txt文件用于声明运行所需依赖,结构紧凑,适合作为标点预测任务的入门参考。目前已有478人学习下载,说明该方案在中文标点恢复方向具有一定关注度。读者可借此了解PaddleNLP标点模型的基本调用方式、推理流程组织与依赖配置思路,并在此基础上替换模型或调整参数,完成自己的文本断句实验。
1. 标点预测这件事,为什么值得用 PaddleNLP 重做一遍
语音转写稿、OCR 识别结果、聊天记录导出、弹幕抓取——这些文本有个共同点:没有标点。一整段几百字连在一起,模型读起来费劲,人读起来更费劲。给预测文本自动补标点,本质是一个序列标注任务:把每个汉字打上「O / 逗号 / 句号 / 问号 / 顿号」这类标签,再按标签把标点插回去。听起来简单,但真上手你会发现,标点预测的坑不在模型结构,而在数据构造、标签对齐和推理时的解码逻辑。
PaddleNLP 在这件事上有天然优势:它内置了 ERNIE 系列预训练模型,TokenClassificationTask直接把序列标注的训练流程封装好了,你不需要自己写 Dataset、不需要手写 CRF 层,改一个配置文件就能跑。我见过太多人用 BERT + 自己搭的 BiLSTM 硬啃,结果卡在 tokenizer 和标签对齐上三天出不来。这篇就按「数据怎么造 → 模型怎么训 → 推理怎么解码 → 坑在哪」的顺序,把基于 PaddleNLP 的标点预测源码拆开讲清楚。适合已经会跑 PaddleNLP 基础 demo、想把它落到真实文本处理流水线里的工程师。
2. 数据从哪来、标签怎么定:标点预测的第一道分水岭
2.1 为什么不能直接用原始文本当训练数据
标点预测的训练数据必须是「无标点文本 → 有标点文本」的配对。但现实里你拿到的往往只有带标点的正常文本,没有现成的无标点版本。常见做法是:拿一批正常文本,用规则把标点全部去掉,生成输入;原始带标点文本作为标签来源。这里有个关键细节——去掉标点后,标签序列的长度必须和输入字符序列严格对齐。
很多人在这里翻车:中文标点去掉后,字符数变了,但标签还是按原文本位置打的,训练时 loss 直接爆炸或者模型学出一堆错位标点。正确做法是逐字符遍历,遇到标点就记录「这个位置应该是什么标点」,遇到非标点就记「O」,保证输入和标签一一对应。
# 构造标点预测训练样本的核心逻辑 # 输入: 带标点的原始句子 # 输出: (无标点字符列表, 对应的标签列表) PUNCT_MAP = { ',': 'COMMA', '。': 'PERIOD', '?': 'QUESTION', '!': 'EXCLAMATION', '、': 'DUNHAO', ':': 'COLON', ';': 'SEMICOLON', '“': 'QUOTE', '”': 'QUOTE', } def build_sample(sentence): chars = [] # 无标点字符序列 labels = [] # 与 chars 等长的标签序列 for ch in sentence: if ch in PUNCT_MAP: # 标点不进入输入序列,而是作为前一个字符的标签 if chars: labels[-1] = PUNCT_MAP[ch] else: chars.append(ch) labels.append('O') return chars, labels这段代码的逻辑是:标点不占输入位置,而是「挂」在它前面那个字符的标签上。比如「今天天气不错,适合出门。」会变成输入['今','天','天','气','不','错','适','合','出','门'],标签['O','O','O','O','O','O','COMMA','O','O','PERIOD']。参数上,PUNCT_MAP决定了你要预测几类标点,类别越多任务越难,建议先从逗号、句号、问号三类起步。
2.2 标签体系设计:别一上来就搞二十类
我见过有人把中文所有标点都列进去,结果模型在顿号、分号、冒号之间反复横跳,F1 惨不忍睹。血泪经验是:标点预测的标签体系要按业务场景裁剪。语音转写场景,逗号和句号占 90% 以上,问号和感叹号少量,顿号、分号几乎不出现。你硬训二十类,等于让模型在长尾类别上浪费容量。
推荐的分阶段策略:
| 阶段 | 标签集合 | 适用场景 | 预期 F1 |
|---|---|---|---|
| 最小可用 | O, COMMA, PERIOD | 语音转写、OCR 后处理 | 0.85+ |
| 标准 | O, COMMA, PERIOD, QUESTION | 通用文本补标点 | 0.80+ |
| 完整 | 上述 + EXCLAMATION, DUNHAO, COLON | 出版级文本处理 | 0.72+ |
先跑最小可用集,确认整条流水线通了,再逐步加类别。每加一类,重新检查混淆矩阵,看新类别是不是在抢其他类别的样本。
2.3 用 PaddleNLP 的 TokenClassificationTask 组织数据
PaddleNLP 的TokenClassificationTask要求数据格式是每行「字符\t标签」,句子之间空行分隔。你需要把上一步的chars和labels写成这种格式:
def write_dataset(samples, out_path): with open(out_path, 'w', encoding='utf-8') as f: for chars, labels in samples: for ch, lb in zip(chars, labels): f.write(f"{ch}\t{lb}\n") f.write("\n") # 空行分隔句子这里有个容易忽略的点:PaddleNLP 的 tokenizer 会对中文按字切分,但如果你用的是 ERNIE 的ernie-3.0-base-zh,它的词表里有些多字词。不过序列标注任务里,tokenizer 默认按字切,is_split_into_words=True时不会把多个字合并。所以你的「字符\t标签」格式能直接对上。写完数据后,用load_dataset读进来,确认len(input_ids) == len(labels),不等就说明有句子被截断或对齐错了。
3. 模型训练:从配置文件到收敛的完整路径
3.1 选 ERNIE 还是 BiLSTM+CRF
PaddleNLP 里做序列标注,最省事的路径是ernie-3.0-base-zh接一个TokenClassificationTask。ERNIE 本身已经在大规模中文语料上预训练过,对标点这种依赖上下文的任务,微调几百步就能出效果。BiLSTM+CRF 是经典方案,但你需要自己搭网络、自己写 CRF 层,而且没有预训练权重,从零训需要的数据量和算力都大得多。
我的建议很直接:有预训练模型就用预训练模型。ERNIE 的注意力机制天然能捕捉「这个位置该不该断句」的上下文信号,比 BiLSTM 的循环结构更高效。CRF 层在标点预测里收益有限,因为标点之间没有强转移约束(逗号后面跟句号也合理),PaddleNLP 的TokenClassificationTask默认不加 CRF,效果已经够用。
3.2 训练脚本与关键参数
PaddleNLP 的训练入口在examples/token_classification/下,核心是run_glue.py的变体。你不需要改模型代码,只需要准备数据、改配置。下面是一个最小训练命令:
python -u run_token_cls.py \ --model_name_or_path ernie-3.0-base-zh \ --train_file ./data/train.txt \ --dev_file ./data/dev.txt \ --max_seq_length 128 \ --learning_rate 3e-5 \ --batch_size 32 \ --epochs 10 \ --save_dir ./output/punct_model \ --logging_steps 50 \ --eval_steps 200 \ --seed 42参数逐个说:max_seq_length设 128 是因为标点预测通常按句处理,单句很少超过 128 字;如果你要处理整段,可以调到 256 或 512,但显存会涨。learning_rate用 3e-5 是 ERNIE 微调的常规起点,太大容易把预训练权重冲垮,太小收敛慢。batch_size32 在 16G 显存上跑 base 模型没问题,如果 OOM 就降到 16。epochs10 是上限,实际看 dev F1,通常 3-5 轮就稳定了。
3.3 训练过程中看什么指标
标点预测不能只看 accuracy。因为「O」标签占绝大多数,模型全预测「O」也能有 80%+ 的准确率,但一个标点都插不对。你要盯的是每个标点类别的 precision、recall、F1。PaddleNLP 的评估会输出 classification report,重点看 COMMA 和 PERIOD 的 F1 是否同步上升。
如果 COMMA 的 recall 高但 precision 低,说明模型到处插逗号;反过来 precision 高 recall 低,说明模型太保守,该断的地方不断。这两种情况的调法不同:前者加负样本(多给一些不该断的句子),后者加正样本(多给一些该断的句子)。我一般会在 dev 集上跑一遍预测,把错例打印出来看,比盯数字有用。
# 训练后快速检查预测结果 from paddlenlp.transformers import ErnieForTokenClassification, ErnieTokenizer model = ErnieForTokenClassification.from_pretrained('./output/punct_model') tokenizer = ErnieTokenizer.from_pretrained('ernie-3.0-base-zh') def predict(text): inputs = tokenizer(text, return_tensors='pd', is_split_into_words=True) logits = model(**inputs)[0] preds = logits.argmax(axis=-1).numpy()[0] # 把标签映射回标点,插到对应字符后面 result = '' for ch, pid in zip(text, preds[1:len(text)+1]): # 跳过 [CLS] result += ch if id2label[pid] != 'O': result += label2punct[id2label[pid]] return result这段推理代码里,preds[1:len(text)+1]是为了跳过[CLS]和[SEP]的位置。如果你发现预测结果整体偏移一位,八成是这里没对齐。id2label和label2punct需要你自己根据训练时的标签映射建。
4. 推理部署:解码逻辑比模型本身更容易出错
4.1 从 logits 到标点序列的解码
模型输出的是每个位置的类别概率,你需要把它转成「在哪个字符后面插哪个标点」。这里有个常见误区:直接取 argmax 就完事。但实际文本里,连续两个逗号、句号后面跟逗号这类情况虽然少见,一旦出现就很扎眼。我的做法是加一层后处理规则:句号后面不再跟逗号,问号后面不再跟句号,连续标点只保留第一个。
def decode_with_rules(chars, pred_ids, id2label): result = [] prev_punct = None for ch, pid in zip(chars, pred_ids): label = id2label[pid] if label != 'O': # 规则1: 句号/问号后不再插逗号 if prev_punct in ('PERIOD', 'QUESTION') and label == 'COMMA': result.append(ch) continue # 规则2: 连续同类标点只保留一个 if prev_punct == label: result.append(ch) continue result.append(ch) result.append(label2punct[label]) prev_punct = label else: result.append(ch) prev_punct = None return ''.join(result)这段解码逻辑不复杂,但能消掉大部分「看起来不对劲」的预测。参数上,prev_punct记录上一个插入的标点类型,用来做规则判断。你可以按业务需求加更多规则,比如引号配对、括号配对,但别加太多,规则太密会跟模型预测打架。
4.2 长文本切分与拼接
实际场景里,你要处理的往往是一整段几百上千字的文本,而模型最大长度只有 512。直接截断会丢信息,正确做法是按句切分或按固定窗口滑动。我一般用滑动窗口:窗口 256 字,步长 128,重叠部分取两次预测中置信度高的那个。
def predict_long_text(text, model, tokenizer, window=256, stride=128): results = [] for start in range(0, len(text), stride): chunk = text[start:start+window] if len(chunk) < 5: # 太短的尾巴不单独预测 break pred = predict(chunk) results.append((start, pred)) # 按位置拼接,重叠区域以先出现的为准 final = '' last_end = 0 for start, pred in results: if start >= last_end: final += pred else: # 重叠部分跳过 overlap = last_end - start final += pred[overlap:] last_end = start + len(pred) return final窗口大小和步长需要按你的文本特点调。语音转写稿句子短,窗口 128 就够;OCR 结果可能整段没标点,窗口可以大一些。步长一般取窗口的一半,保证重叠区足够做拼接判断。
4.3 批量推理的性能优化
如果你要处理几万条文本,逐条推理太慢。PaddleNLP 支持 batch 推理,把多条文本 padding 到同一长度,一次前向。注意 padding 的位置要在计算 loss 和 argmax 时 mask 掉,否则会预测出乱七八糟的标点。
def batch_predict(texts, model, tokenizer, batch_size=16): all_results = [] for i in range(0, len(texts), batch_size): batch = texts[i:i+batch_size] inputs = tokenizer(batch, padding=True, return_tensors='pd', is_split_into_words=True) logits = model(**inputs)[0] preds = logits.argmax(axis=-1).numpy() for j, text in enumerate(batch): # 只取有效长度部分 valid_len = len(text) all_results.append(decode_with_rules( text, preds[j][1:valid_len+1], id2label)) return all_resultspadding=True会自动补齐到 batch 内最长,valid_len用来截掉 padding 部分的预测。batch_size 根据显存调,16 或 32 都行。如果文本长度差异大,可以按长度排序后再分批,减少 padding 浪费。
5. 避坑与排查:标点预测里最容易翻车的五个地方
5.1 现象:训练 loss 正常下降,但推理结果全是「O」
原因:标签分布极度不均衡,「O」占 90% 以上,模型学到「全预测 O」就能拿到低 loss。解决:在数据构造时对非 O 标签做上采样,或者用class_weight给非 O 类别更高权重。PaddleNLP 的TokenClassificationTask支持传class_weight,我一般给非 O 类别 5-10 倍权重。
5.2 现象:预测的标点位置整体偏移一位
原因:tokenizer 返回的input_ids包含[CLS]和[SEP],解码时没跳过。解决:确认preds[1:len(text)+1]这个切片,[CLS]在位置 0,[SEP]在最后。如果你用的是is_split_into_words=True,还要检查 tokenizer 有没有把某些字拆成 subword。
5.3 现象:逗号预测过多,几乎每个短语后面都插逗号
原因:训练数据里逗号密度过高,或者 COMMA 类别的 recall 权重设太大。解决:统计训练集里逗号出现的平均间隔,如果每 5 个字就一个逗号,说明数据本身有问题。正常中文文本逗号间隔在 10-15 字左右。另外检查class_weight,别给 COMMA 单独加太高。
5.4 现象:长文本推理时后半段标点质量明显下降
原因:滑动窗口拼接时重叠区处理不当,或者窗口太大导致模型注意力分散。解决:把窗口从 512 降到 256,步长设为 128,确保每个字符至少被预测两次。拼接时以第一次出现的预测为准,避免边界处标点重复。
5.5 现象:模型在测试集上 F1 很高,上线后用户反馈「标点奇怪」
原因:测试集和真实场景的文本分布不一致。测试集可能是新闻语料,线上是口语化聊天记录。解决:从线上日志里采样一批真实文本,人工标注 200-500 条作为验证集,重新评估。如果差距大,用线上数据做增量微调,学习率调小到 1e-5。
6. 进阶技巧:用置信度阈值和规则兜底把可用性再提一档
模型输出的是概率,argmax 只是取了最大值。但有些位置模型本身就不确定,比如「今天天气不错适合出门」——「错」后面到底该不该断,模型可能给 COMMA 0.4、O 0.35、PERIOD 0.25。这种位置硬插标点,错了比不插更难看。我的做法是设一个置信度阈值:只有最高概率超过 0.6 才插标点,否则保持 O。这个阈值在验证集上调,看 precision 和 recall 的平衡点。
def decode_with_threshold(chars, probs, id2label, threshold=0.6): result = [] for ch, prob in zip(chars, probs): max_id = prob.argmax() max_prob = prob[max_id] if max_prob >= threshold and id2label[max_id] != 'O': result.append(ch) result.append(label2punct[id2label[max_id]]) else: result.append(ch) return ''.join(result)阈值从 0.5 开始试,每 0.05 一档,看 dev 集上 F1 的变化。通常 0.55-0.65 之间效果最好。低于 0.5 会插太多错标点,高于 0.7 会漏掉该断的地方。
另一个技巧是规则兜底:模型预测完之后,用正则扫一遍,把明显不合理的标点组合修掉。比如连续两个句号、逗号后面紧跟句号、引号不配对。这些规则不需要多,十条以内就能覆盖 90% 的异常情况。规则和模型的关系是:模型负责「哪里该断」,规则负责「断得对不对」。
| 规则 | 正则/逻辑 | 处理方式 |
|---|---|---|
| 连续句号 | 。。+ | 合并为一个 |
| 逗号后跟句号 | ,。 | 删掉逗号 |
| 句号后跟逗号 | 。, | 删掉逗号 |
| 问号后跟句号 | ?。 | 删掉句号 |
| 开头标点 | ^[,。?!] | 删掉 |
最后说个我自己的习惯:每次训完模型,我会随机抽 50 条线上文本,人工过一遍预测结果,把错例按「多插、漏插、插错类型」分类统计。如果某一类错误超过 20%,就针对性地补数据或调阈值。标点预测没有一步到位的方案,都是靠这种小步迭代磨出来的。希望帮到你。
本文还有配套的精品资源,点击获取