Skip to content

Alidadei/LearnPostTrain

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

25 Commits
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

LearnPostTrain

面向有限显存的大语言模型后训练实验项目。当前版本以 GRPO + LoRA 为主,重点验证在单张 GPU 上进行可复现、可诊断的 Countdown 任务训练。

English README

项目定位

本仓库是 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 结果直接比较。

Hook 设计

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、模型版本和数据划分单独宣称有效:

  1. Warmup + cosine 学习率:全量训练使用 8e-6 → 1e-6 的 600-step warmup-cosine;相比固定 1e-5,验证集峰值从约 0.45 提升到 0.54,canonical test-1000 为 0.478。
  2. 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)可以运行,但现有消融没有超过固定系数。
  3. 更大有效 batch:LoRA 显存较低,主实验使用 batch_size=25632 个问题 × 8 个回答;全量实验使用 128。LoRA 的优势部分来自省下显存后可以看到更多训练题,而不只是可训练参数更少。
  4. Completion-only logits:只为 completion 位置计算 policy logits,避免为 prompt/padding 位置构造完整词表 logits;数值语义回归与旧路径一致。
  5. Selective activation checkpointing:当前采用 selectivesegments=2。局部 3B LoRA benchmark 的 update 中值比旧路径快约 17.7%,同 seed 端到端 1-step 总时长下降 6.3%;600-step 墙钟收益是外推值,不应当当作完整复测结果。
  6. 减少无效日志与显存碎片:正式配置关闭高频进度输出,训练代码默认使用 expandable_segments;24 GB 全量配置提供 optimizer CPU offload。当前 micro_batch_size=2 是 24 GB 下已验证的稳妥值,更大的 micro-batch 没有形成可复现的端到端收益。
  7. 可续训 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.yaml

24 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.py

GPU 评估和 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

About

A reproducible, low-memory post-training framework for GRPO, OPD, and DPO with LoRA on consumer GPUs.

Resources

Code of conduct

Contributing

Security policy

Stars

Watchers

Forks

Releases

Packages

Contributors

Languages