Skip to content

Repository files navigation

Predictive Text

轻量级中文预测性文本输入模型,基于 Transformer 架构,支持 ONNX 部署。

这是为 https://github.com/ximeiorg/Xime 项目准备的联想词预测模型。

预训练模型

模型 大小 下载
base 23M 参数 ModelScope
small 10M 参数 ModelScope

功能特性

  • 支持多种模型尺寸(small/base/large)
  • 基于 label.txt 常用汉字表的词表剪枝,小参数模型也能有效收敛
  • 训练时自动忽略句中逗号/顿号的 loss,避免模型偏好预测标点
  • ONNX 导出与量化(INT8)
  • 动态序列长度,无需填充

快速开始

# 安装依赖
uv sync

# 1. 准备数据:BPE 训练 → 频率分析 → label.txt 过滤剪枝 → tokenize
uv run python scripts/analyze_and_prune_vocab.py

# 2. 训练模型
uv run src/train.py --model-size small --use-prepared-data

# 3. 导出 ONNX (float32 + INT8)
uv run scripts/export_onnx.py --model-size small

对齐模型

对齐脚本

 uv run python scripts/compare_pt_onnx.py

联想预测示例

使用训练好的模型进行联想预测:

small (10M)

  今天天气 -> 晴 | 预 | 很 | 真 | 怎么样
  我们一起去 -> 看 | 看看 | 找 | 公园 | 做
  我觉得 -> 这 | 你 | 自己 | 我 | 这个
  明天早上 -> 点 | 起来 | 就 | 起 | 来
  我在北京 -> 做 | 吃 | 旅行 | 读 | 的

base (23M)

  今天天气 -> 之 | 下 | 是 | 明 | 天
  我们一起去 -> 打 | 工 | 是 | 必然 | 的
  我觉得 -> 很 | 明 | 智 | 是 | 正确
  明天早上 -> 上 | 班 | 了 | 是 | 真的
  我在北京 -> 都 | 练 | 马 | 师 | 的
  正在吃饭 -> 的时候 | , | 我 | 也 | 在

模型对比

指标 small (10M) base (23M) 提升
Perplexity 96.84 31.09 -68%
Top-1 29.78% 39.02% +9.2pp
Top-3 40.74% 53.40% +12.7pp
Top-5 45.76% 60.02% +14.3pp
Top-10 52.36% 67.54% +15.2pp
评分 23.3/100 40.5/100 +17.2

base 相对 small:参数量翻倍(10M → 23M),Top-1 提升 9 个百分点,Top-5 提升 14 个百分点。
评测基于 5000 条验证集采样,详见 output/benchmark/compare_base_rope_vs_no_rope.md

词表剪枝流程

BPE 词表(~12000)对中文会产生大量低效的字节级碎片 token。 使用 label.txt 常见汉字表 + 频率分析做剪枝:

  1. 训练标准 BPE tokenizer(12000 词表)
  2. 采样 30% 语料统计每个 token 频次
  3. 分类:
    • 单字在 label.txt 中 → 保留
    • 单字不在 label.txt 中 → 丢弃(覆盖率 < 0.2%)
    • 多字组合 → 按频率排序保留 top-N
  4. 输出 data/vocab.json(~8000 词表)+ data/train.bin + data/val.bin

标点 loss 控制

训练时 WikipediaDataset 自动检测词表中的标点符号(,、;:。!?·~—…()【】《》""''「」『』 等), 将其 loss 乘以 0.1 降低权重,减少标点对联想词的干扰。

模型配置

名称 vocab dim heads layers ffn seq 参数量 float32
small 8000 384 4 4 1536 64 10.2M 39 MB
base 8000 512 8 6 2048 128 23.0M 92 MB
large 8000 768 12 8 3072 128 62.9M 240 MB

目录结构

├── src/
│   ├── train.py          # 训练脚本
│   ├── inference.py      # 推理脚本
│   ├── inference_candidate.py  # 联想词推理
│   ├── config.py         # 配置定义
│   └── model/            # 模型实现
├── scripts/
│   ├── analyze_and_prune_vocab.py  # 词表分析+剪枝 (核心)
│   ├── prepare_from_cleaned.py     # 数据准备 (备选)
│   ├── export_onnx.py              # ONNX 导出 (float32 + INT8 量化)
│   ├── test_onnx_inference.py      # ONNX 测试
│   └── verify_model.py             # 模型验证
├── data/                 # 数据目录 (vocab.json / train.bin / val.bin)
├── output/               # 训练输出
├── label.txt             # 常用汉字表 (用于词表剪枝)
└── mobile/               # ONNX 模型

文档

License

CC BY-NC-SA

About

一个专门应用于手机输入法的智能联想模型

Resources

Stars

Watchers

Forks

Releases

Packages

Contributors

Languages