面向有限显存的大语言模型后训练实验项目。当前版本以 GRPO + LoRA 为主,重点验证在单张 GPU 上进行可复现、可诊断的 Countdown 任务训练。
本仓库是 policy-gradient/GRPO-Zero 的独立研究分支,保留了从零实现的 PyTorch GRPO 基线,并加入了面向实验复现的训练、评估和诊断组件。
当前发布版本明确限定在以下范围:
- 基础模型:Qwen2.5-3B-Instruct;
- 数据集:Countdown-Tasks-3to4;
- 算法:GRPO;
- 训练方式:单 GPU,全量训练或纯 PyTorch LoRA;
- 评估:canonical-key-disjoint 划分及其顺序鲁棒性检查。
这不是一个已经支持任意模型和任意后训练算法的通用框架。OPD、DPO 和更多模型后端属于后续扩展方向,目前不应当理解为已经实现并经过本仓库验证。
train.py:原始风格的最小 GRPO 基线训练器;train_v2.py:基于 Hook 的实验训练器,支持学习率/熵调度、计时、安全告警、早停、互信息诊断和 rollout 过滤;hooks.py:可组合的训练生命周期 Hook;- 纯 PyTorch LoRA、只在 completion 位置计算 policy logits、LoRA checkpoint 续训,以及可选的
torch.compile推理编译; - canonical-key 划分生成器和验证脚本;
- 不加载大模型即可运行的 CPU 回归测试,以及需要 GPU 的 smoke test 和 benchmark。
模型、数据、划分、checkpoint 和日志路径由 YAML 配置指定,不绑定某一台开发机。
配置文件统一采用“算法 + 训练范围 + 模型 + 关键超参 + 用途”的命名方式,例如:
grpo_lora_qwen25_r8_lr5e-5_cos600_ec3e-3.yaml。
配置清单、字段约定和推荐组合见 CONFIG/README.md。当前最值得优先复现的配置是:
| 配置 | 用途 | 已记录结果 |
|---|---|---|
CONFIG/grpo_lora_qwen25_r8_lr5e-5_cos600_ec3e-3.yaml |
当前推荐 LoRA 配方 | v2 canonical test-1000:0.624(step 600) |
CONFIG/grpo_full_qwen25_lr8e-6_cos600_ec1e-3.yaml |
当前推荐全量 GRPO 配方 | v2 canonical test-1000:0.478;val 峰值 0.54 |
CONFIG/grpo_lora_qwen25_r8_lr6e-5_cos600_ec2p5e-3_stage1_300step.yaml + 续跑配置 |
更高 LR 的对照实验 | v2 canonical test-1000:0.527,不是当前最佳 |
旧的 config_* sweep 文件是本机实验遗留物,默认被 .gitignore 忽略;发布入口使用上述语义化配置,不再依赖 config_expA 这类难以判断用途的名称。
发布版使用 data/splits/countdown_v2_seed20260701.json,生成种子为 20260701,文件 SHA-256 为 64bb0fd9f1733f449dcea0dc14d1c4a643ce5294ade9b29dec7aba2e1baaedd2。
划分不是按原始行随机切分,而是先按 canonical key 分组:
canonical_key = (target, sorted(nums))
这样,同一个数学问题仅仅改变输入数字顺序时,仍然属于同一组,不会同时出现在 train、validation 和 test 中。生成器先用固定 seed 打乱 canonical keys,然后分配:
- test:1000 个 canonical keys,每个 key 取一个代表行;
- validation:500 个 canonical keys,每个 key 取一个代表行;
- train:其余 448,070 个 canonical keys 的全部排列,共 488,728 行。
清单总计 490,364 行、449,570 个 canonical keys,并额外从 test keys 构造 200 个顺序鲁棒性问题(每题最多 8 个排列)。tests/verify_countdown_split_v2.py 会同时检查行索引不重叠和 canonical key 不重叠。
旧的 legacy 配置使用尾部切分或随机切分,可能发生排列泄漏,只能用于复现历史现象,不能与 canonical-disjoint 结果直接比较。
train_v2.py 使用 hooks.py 中的可组合 Hook,把调度、诊断和过滤等横切逻辑从核心 rollout/update 循环中分离出来。单个训练 step 的生命周期为:
on_step_start -> rollout -> on_after_rollout -> update
-> on_after_update -> eval/checkpoint -> on_step_end
内置 Hook 覆盖学习率和熵调度、计时与安全告警、早停、互信息(MI)诊断以及基于 reward 方差的 rollout 过滤。配置文件只能选择并参数化已经写好的行为,不能仅靠 YAML 定义新的 Python 训练逻辑。
Hook 适合放置调度器、诊断器、过滤器等横切功能;训练主循环和算法特有的目标函数仍保持显式实现。未来加入 OPD 或 DPO 时,应使用独立 trainer/backend,而不是把核心 loss 隐藏在通用 Hook 中。当前模型后端也是 Qwen2 风格的专用实现,Hook 系统不会自动把它变成架构无关的后端。
这些技巧改变的是训练效率或稳定性,不能脱离配置、seed、模型版本和数据划分单独宣称有效:
- Warmup + cosine 学习率:全量训练使用
8e-6 → 1e-6的 600-step warmup-cosine;相比固定1e-5,验证集峰值从约 0.45 提升到 0.54,canonical test-1000 为 0.478。 - LoRA 专用学习率与熵系数:rank=8 LoRA 只训练约 15M 参数(约 0.48%)。
5e-5 + entropy_coeff=0.001很快熵坍缩,0.005又会把熵推到约 11.4 的近均匀随机策略;5e-5 + 0.003是当前记录中兼顾稳定性和效果的组合。entropy_schedule的分段衰减(step 0/100/300 对应0.001/0.0005/0.0001)可以运行,但现有消融没有超过固定系数。 - 更大有效 batch:LoRA 显存较低,主实验使用
batch_size=256、32个问题 ×8个回答;全量实验使用128。LoRA 的优势部分来自省下显存后可以看到更多训练题,而不只是可训练参数更少。 - Completion-only logits:只为 completion 位置计算 policy logits,避免为 prompt/padding 位置构造完整词表 logits;数值语义回归与旧路径一致。
- Selective activation checkpointing:当前采用
selective、segments=2。局部 3B LoRA benchmark 的 update 中值比旧路径快约 17.7%,同 seed 端到端 1-step 总时长下降 6.3%;600-step 墙钟收益是外推值,不应当当作完整复测结果。 - 减少无效日志与显存碎片:正式配置关闭高频进度输出,训练代码默认使用
expandable_segments;24 GB 全量配置提供 optimizer CPU offload。当前micro_batch_size=2是 24 GB 下已验证的稳妥值,更大的 micro-batch 没有形成可复现的端到端收益。 - 可续训 checkpoint:新的 LoRA checkpoint 保存 adapter、optimizer、RNG 和 DataLoader 位置,可从中断点继续;旧的 adapter-only checkpoint 首次恢复时无法重建历史 optimizer 状态。
MI 检测和 rollout filter 仍是可选诊断/实验项。当前 filter 实验的 test accuracy 低于未过滤版本(约 35.6% vs 38.6%),所以规范 LoRA 配置默认关闭它们。
下表只列出已保存记录中的结果;不同 evaluator、数据划分或采样设置不能混为同一条曲线。
| 配置 | 训练范围 | 结果 | 结论 |
|---|---|---|---|
grpo_full_qwen25_lr1e-5_legacy.yaml / canonical 复现 |
全量,600 步以内 | canonical val 约 0.45,entropy 约 0.06 | 学习率过激,出现熵锁死/恶性坍缩 |
grpo_full_qwen25_lr8e-6_cos600_ec1e-3.yaml |
全量,600 步 | val 峰值 0.54;v2 test-1000 0.478;legacy test-500 0.486 | 当前全量最佳记录 |
grpo_lora_qwen25_r8_lr5e-5_cos600_ec3e-3.yaml |
LoRA,600 步 | v2 test-1000 0.624;val 0.61;约 14–18 GB 显存 | 当前整体最佳记录;不是与全量完全等预算的对照 |
grpo_lora_qwen25_r8_lr6e-5_cos600_ec2p5e-3_* |
LoRA,600 步 | v2 test-1000 0.527;val 峰值约 0.512 | 更高 LR 的可复现实验,但不超过 5e-5 配方 |
LoRA 最佳记录使用了全量实验两倍的 rollout 规模(batch 256 vs 128),因此结果应表述为“低显存 LoRA 配方在本实验预算下取得更高准确率”,而不是简单归因于 LoRA 本身。
- 推荐 Linux 或 WSL2;
- Python 3.11 或更高版本;
- 实际 Qwen 训练需要 CUDA GPU,单元测试可在 CPU 上运行;
- 推荐使用
uv管理环境。
模型和 Countdown 数据集是外部资源,其许可证和使用条款以各自提供方为准,本仓库不重新分发这些权重或数据。
安装环境并准备外部模型/数据:
uv sync --group dev
git clone https://huggingface.co/Qwen/Qwen2.5-3B-Instruct
git clone https://huggingface.co/datasets/Jiayi-Pan/Countdown-Tasks-3to4生成并验证 canonical split:
uv run python scripts/create_countdown_split_v2.py \
--data-path Countdown-Tasks-3to4 \
--output data/splits/countdown_v2_seed20260701.json
uv run python tests/verify_countdown_split_v2.py运行当前推荐配置:
uv run python train_v2.py \
--config CONFIG/grpo_lora_qwen25_r8_lr5e-5_cos600_ec3e-3.yaml运行当前推荐的全量配置:
uv run python train_v2.py \
--config CONFIG/grpo_full_qwen25_lr8e-6_cos600_ec1e-3.yaml24 GB legacy offload 示例、LoRA 单 step smoke 和 timing 配置见 CONFIG/README.md。
默认测试不会加载大模型:
uv run pytest -q常用的定向检查:
uv run python tests/verify_lora_wiring.py
uv run python tests/verify_lora_step1.py
uv run python tests/verify_training_speedup_semantics.py
uv run python tests/test_lora_checkpoint_resume.py
uv run python tests/verify_countdown_split_v2.pyGPU 评估和 benchmark 需要外部模型、数据集及 checkpoint 路径;环境变量和命令行路径说明见 docs/reproducibility.md。
README 中的准确率等数字属于实验结果,不是对任何硬件或随机种子的保证。要复现实验,至少应记录 commit、配置、seed、模型和数据集版本、split hash、样本数、硬件、采样设置以及完整命令。Checkpoint 和 TensorBoard 日志不纳入 Git。
当前实现依赖 Qwen2/Qwen2.5 风格的模型结构和 tokenizer/输出约定。只修改配置中的模型路径,并不能保证直接运行 Qwen3、Llama、Mistral 或 Gemma;其他架构需要新增后端、输出解析和训练/评估验证。
| 方向 | 当前状态 |
|---|---|
| GRPO + Qwen2.5 | 已实现并有配置/测试 |
| LoRA 低显存训练 | 已实现 |
| Hook 调度与诊断体系 | 已实现 |
| OPD / DPO | 规划中,尚未作为当前版本算法发布 |
| Qwen3/其他模型架构 | 需要新增后端,当前未承诺开箱即用 |
本项目采用 Apache License 2.0。上游项目及第三方工作见 NOTICE 和源 README 历史说明;本分支的变更记录见 CHANGELOG.md。