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.ptcheckpoints/potato.ptcheckpoints/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.pyConfig 数据类中。修改行为前务必先阅读该文件——许多功能存在但默认关闭:

  • 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=Falsedynamic_experts=Falsemoe_router_noise=0.0 — 渐进式缩放 / 动态专家分配 / 路由噪声均为关闭的预留能力。

Rust 引擎(可选加速)

cargo build --release --lib
copy target\release\potato_engine.dll potato_lm\potato_engine.pyd   # Windows

提供的内核:parallel_matmul_f32fused_moe_forward/backwardsparse_moe_topkscaled_dot_product_attentiongelu/silu_fusedrms_norm_fusedcross_entropy_loss_fused

注意事项:

  • 引擎缺失或 dtype/shape 不匹配时会静默回退到 PyTorch,先用小 shape 验证。
  • 训练陷阱Attentionmodel.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,使旧检查点失效,需要重新训练。
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Papers for Xuchen818/Potato