YAML Metadata Warning:empty or missing yaml metadata in repo card
Check out the documentation for more information.
Potato v4
面向 CPU 的轻量级 MoE(Mixture-of-Experts)训练平台。模型以真实 MNIST 手写数字图像为输入,通过"思考 → 调用工具 → 观察结果 → 验证 → 作答"的 agentic 轨迹完成整数算术(+ - *)、大小比较(> < =)和两步链式计算任务。
该项目是完整的训练与推理工具链:视觉编码器预训练、GLM 式空白填充预训练、三阶段 SFT、GRPO 强化学习,以及自一致性投票评估。所有计算均在 CPU 上运行(PyTorch + 可选 Rust 内核加速),不依赖 GPU。
详细的设计说明与实验记录见 potato_technical_report.pdf。
技术要点
- 多模态 MoE 架构:CNN 视觉编码器(MNIST 图像 → 40 个图像 token)+ 4 层 Transformer,MoE 层含 64 个专家、每 token 激活 Top-8 [2][3]。
- 字符级 tokenizer:19 个特殊 token + 92 个字符,共 111 词表,支持
think / step / tool_call / observation / verify / answer完整轨迹标记。 - 多阶段训练:视觉预训练 → GLM 空白填充预训练 [4] → SFT(高学习率 → 低学习率 → 自蒸馏)→ GRPO 强化学习 [7](自适应 KL)。
- Agentic 工具调用:模型自主调用
calc(算术求值,正则白名单,无eval)和cmp(比较)两个安全工具 [6]。 - Rust 加速引擎:基于 PyO3 的 CPU 内核(MoE 前向/反向、注意力、RMSNorm、交叉熵等),缺失或形状不匹配时静默回退到 PyTorch。
- 纯 Python 依赖:仅需 PyTorch、NumPy、Pillow;MNIST 数据首次运行时自动下载,失败则回退到合成数字。
环境要求
| 组件 | 要求 |
|---|---|
| Python | 3.10+(仓库在 3.13 上验证) |
| PyTorch | ≥ 2.0(CPU 版即可,torch 官方 CPU wheel 建议使用 --index-url https://download.pytorch.org/whl/cpu 安装) |
| NumPy / Pillow | 标准 pip 安装 |
| Rust 工具链(可选) | 构建 potato_engine.pyd 加速内核需要;不装则自动回退 PyTorch |
| MSVC Build Tools(可选) | 供 moe_c_extension.py 的 C++ 路径使用(load_inline JIT 编译,需 vcvarsall.bat) |
pip install torch numpy pillow
快速开始
首次运行会检查并自动下载 MNIST 数据集(存至 data/mnist/),无需手动准备数据。
完整训练流水线
python train_full.py
按顺序执行:GLM 预训练 → SFT 三阶段 → GRPO 强化学习。各阶段检查点分别写入 checkpoints/potato_glm.pt、checkpoints/potato.pt、checkpoints/potato_grpo.pt。
分阶段训练(每阶段可独立运行)
python -u -m potato_lm.pretrain_vision # 1) CNN 视觉编码器(MNIST 数字识别)→ vision_conv.pt
python -u -m potato_lm.glm_pretrain # 2) GLM 空白填充预训练 → potato_glm.pt
python -u -m potato_lm.train # 3) SFT(高LR → 低LR → 自蒸馏)→ potato.pt
python -u -m potato_lm.grpo # 4) GRPO 强化学习 → potato_grpo.pt
指定阶段与参数
python train_full.py --glm-only # 只跑 GLM 预训练
python train_full.py --sft-only # 只跑 SFT
python train_full.py --grpo-only # 只跑 GRPO
python train_full.py --resume # 从已有检查点继续
python train_full.py --threads 8 --epochs 10 --batch-size 32
评估与演示
py -3.13 eval_checkpoint.py # potato.pt 上 48 样本评估
py -3.13 eval_grpo.py # potato_grpo.pt 上 48 样本评估
py -3.13 eval_sc.py --k 8 --temp 0.8 # 自一致性投票评估 [8]
py -3.13 demo.py # 24 样本交互式演示(有 grpo 检查点则用 grpo,否则用 potato)
py -3.13 smoke_test.py # 端到端冒烟测试:5 epochs + 4 个 agentic 评估
py -3.13 _test_rust_autograd.py # 校验 Rust MoE 梯度与 PyTorch autograd 一致
训练配置
所有超参数集中在 potato_lm/config.py 的 Config 数据类中。修改行为前务必先阅读该文件——许多功能存在但默认关闭:
use_glm_pretrain=False— GLM 预训练默认不启用use_self_distill=False— 自蒸馏默认关闭(此前实验显示会降低准确率)external_data_ratio=0.0— 外部 HF 数据集混合默认关闭(pickle 样本的 token ID 与当前 tokenizer 不兼容,启用需重新分词)phase2_aux_loss_weight=0.0— 有意关闭,使 loss 可以降到 0.001 以下resume=True— 默认从已有potato.pt热启动
另注意与直觉相反、声明启用但实际未接线或静默的开关:
use_self_learning=True/use_online_dpo=True— 自我学习(RAG + 在线 DPO)当前流水线并未接线:generate.py仅当调用方显式传入self_learner参数时才触发,现有train_full.py/train.py/ 各 eval 脚本均未传入,属预留能力。knowledge.py为其 RAG 后端。use_simd=True/use_rust_engine=True— 加速内核默认开启;use_simd启动时会尝试加载 C++ 扩展(moe_c_extension._get_module(),train.py:493),依赖 MSVC Build Tools,失败仅静默提示。code_data_ratio=0.0— 代码语法骨架数据默认关闭(run 13 曾用于教模型续写,见code_data.py)。progressive_scale=False、dynamic_experts=False、moe_router_noise=0.0— 渐进式缩放 / 动态专家分配 / 路由噪声均为关闭的预留能力。
Rust 引擎(可选加速)
cargo build --release --lib
copy target\release\potato_engine.dll potato_lm\potato_engine.pyd # Windows
提供的内核:parallel_matmul_f32、fused_moe_forward/backward、sparse_moe_topk、scaled_dot_product_attention、gelu/silu_fused、rms_norm_fused、cross_entropy_loss_fused。
注意事项:
- 引擎缺失或 dtype/shape 不匹配时会静默回退到 PyTorch,先用小 shape 验证。
- 训练陷阱:
Attention(model.py:54)直接使用F.scaled_dot_product_attention。切勿在训练中换成 Rust 注意力内核(engine_bridge.fast_attention)——它内部调用.detach(),会静默切断 Q/K/V 权重梯度。Rust bridge 函数仅用于推理;训练路径下唯一安全的 Rust 代码是rust_autograd.RustMoEFunction。 - 修改 Rust 源码后可用
_test_rust_autograd.py验证——它直接从target/release/加载新 DLL,无需先部署.pyd。
项目结构
potato_lm/ # Python 包(全部训练逻辑)
pretrain_vision.py # 阶段 1:CNN 视觉预训练
glm_pretrain.py # 阶段 2:GLM 空白填充预训练
train.py # 阶段 3:三阶段 SFT
grpo.py # 阶段 4:GRPO 强化学习
model.py # 模型架构(VisionEncoder + MoE Transformer)
moe.py # MoE 层(Rust/C++/Python 三级路径)
rust_autograd.py # Rust MoE 的 autograd 封装
engine_bridge.py # Rust 引擎桥接(自动回退)
dataset.py # 数据生成与 MNIST 图像合成
mnist.py # MNIST 下载/解析(无 torchvision 依赖,失败回退合成数字)
tokenizer.py # 字符级 tokenizer
generate.py # agentic 推理循环(支持可选 self_learner 参数)
tools.py # 安全工具(calc / cmp)
moe_c_extension.py # C++ 扩展:MoE 前向(JIT 编译,load_inline,最佳努力加载)
simd_kernels.py # SIMD/MKL-DNN 线程配置 + numpy 向量化 MoE 前向
self_learn.py # 自我学习(RAG + 在线 DPO)——预留模块,默认未接线
knowledge.py # RAG 知识库(纯 Python TF-IDF 余弦检索)
code_data.py # 代码语法骨架数据生成(默认关闭)
config.py # 全部超参数
src/ # Rust crate potato-engine(PyO3)
train_full.py # 全流程编排入口
*.py # 评估 / 演示 / 冒烟测试脚本
checkpoints/ # 模型权重(potato.pt / potato_grpo.pt 等)
external_data/ # 预下载的 HF 数据集(当前未启用)
参考文献
| # | 技术 | 文献 |
|---|---|---|
| [1] | Transformer 基础架构 | Vaswani et al., Attention Is All You Need, NeurIPS 2017. https://arxiv.org/abs/1706.03762 |
| [2] | 稀疏门控 MoE | Shazeer et al., Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer, ICLR 2017. https://arxiv.org/abs/1701.06538 |
| [3] | MoE 大规模扩展 | Fedus et al., Switch Transformers: Scaling to Trillion Parameter Models with Simple and Efficient Sparsity, JMLR 2022. https://arxiv.org/abs/2101.03961 |
| [4] | GLM 空白填充预训练 | Du et al., GLM: General Language Model Pretraining with Autoregressive Blank Infilling, ACL 2022. https://arxiv.org/abs/2103.10360 |
| [5] | 思维链轨迹 | Wei et al., Chain-of-Thought Prompting Elicits Reasoning in Large Language Models, NeurIPS 2022. https://arxiv.org/abs/2201.11903 |
| [6] | 工具调用与观察循环 | Yao et al., ReAct: Synergizing Reasoning and Acting in Language Models, ICLR 2023. https://arxiv.org/abs/2210.03629 |
| [7] | GRPO 强化学习 | Shao et al., DeepSeekMath: Pushing the Limits of Mathematical Reasoning in Open Language Models, 2024. https://arxiv.org/abs/2402.03300 |
| [8] | 自一致性投票 | Wang et al., Self-Consistency Improves Chain of Thought Reasoning in Language Models, ICLR 2023. https://arxiv.org/abs/2203.11171 |
| [9] | 在线偏好优化 | Rafailov et al., Direct Preference Optimization: Your Language Model is Secretly a Reward Model, NeurIPS 2023. https://arxiv.org/abs/2305.18290 |
| [10] | MNIST 数据集 | LeCun et al., Gradient-Based Learning Applied to Document Recognition, Proc. IEEE 1998. https://ieeexplore.ieee.org/document/726791 |
常见问题
- MNIST 下载失败? 运行会自动重试,失败时回退到程序生成的合成数字图像,训练仍可继续。
- 训练很慢? 可通过
--threads N控制线程数;确认potato_lm/potato_engine.pyd存在且与当前 Rust 源码一致(_test_rust_autograd.py可验证)。use_bf16_autocast=True在无 AVX512-BF16 的 CPU 上反而更慢。 - 加载旧检查点报词表错误? 早期检查点是 109 词表(不含
?token)。eval_sc.py会自动将 tokenizer 适配到检查点词表;cell_token_ids必须取自检查点本身,不要重新推导。 - 更换数据集渲染方式后准确率下降? 数据集生成(
dataset.py)是渲染式的,改变字形/停止符布局会改变cell_token_ids,使旧检查点失效,需要重新训练。