本作业实现语言模型的后训练(Post-Training)对齐流程,包括:
- SFT(Supervised Fine-Tuning) — 在推理轨迹上进行监督微调
- Expert Iteration — 基于策略生成 + 奖励筛选的迭代式训练
- GRPO(Group Relative Policy Optimization) — 基于组归一化奖励的策略梯度优化
- DPO(Direct Preference Optimization) — 直接偏好优化(可选补充作业)
cs336-a5/
├── cs336_alignment/ # 主包
│ ├── __init__.py
│ ├── SFT_util/ # SFT 工具函数
│ │ ├── __init__.py
│ │ ├── compute_entropy_func.py # 熵、log-probs、masked normalize
│ │ ├── tokenize_func.py # prompt+output tokenization
│ │ ├── train_step.py # SFT micro-batch train step + log_generations
│ │ └── model_test.py # 模型评估
│ ├── prompts/ # Prompt 模板
│ │ ├── r1_zero.prompt
│ │ ├── alpaca_sft.prompt
│ │ ├── question_only.prompt
│ │ └── zero_shot_system_prompt.prompt
│ ├── tests/ # 单元测试(来自 handout 的 snapshot 测试)
│ │ ├── __init__.py
│ │ ├── adapters.py # 测试适配器(连接实现与测试)
│ │ ├── test_sft.py # SFT 相关测试
│ │ ├── test_grpo.py # GRPO 相关测试
│ │ ├── test_data.py # 数据加载测试
│ │ ├── test_metrics.py # 评价指标测试
│ │ └── test_dpo.py # DPO 测试
│ ├── run_sft.py # SFT 训练主脚本
│ ├── run_grpo.py # GRPO/Expert Iteration 训练主脚本
│ ├── drgrpo_grader.py # MATH 打分函数(格式 + 答案判分)
│ └── plot_sft_curves.py # 训练曲线可视化
├── data/ # 数据集
│ ├── alpaca_eval_hf/ # AlpacaEval 评估基准
│ ├── gsm8k/ # GSM8K 数学推理基准
│ ├── hh/ # Anthropic HH-RLHF 偏好数据
│ ├── mmlu_hf/ # MMLU 多任务语言理解基准
│ └── simple_safety_tests/ # 简单安全测试集
├── models/ # 预训练模型权重(gitignored)
├── outputs/ # 训练输出(checkpoints、日志)
├── logs/ # 运行日志
├── slurm/ # SLURM 作业脚本
│ ├── run_sft.slurm
│ └── run_grpo.slurm
├── pyproject.toml # 项目配置 + 依赖
├── CHANGELOG.md
└── README.md
本作业使用 uv 管理依赖和 Python 环境(Python 3.12)。
# 安装 uv(如果还没有)
apt update && apt install curl -y
curl -LsSf https://astral.sh/uv/install.sh | sh
source $HOME/.local/bin/env
# 安装项目依赖
# uv sync --no-install-package flash-attn
uv sync这个 Flash Attention 确实很难搞。确定你当前环境的PyTorch版本、CUDA版本、Python版本
python -c "import torch; print(f'Torch: {torch.__version__}, CUDA: {torch.version.cuda}')"
nvcc --version # 检查系统CUDA编译器版本
python --version
python -c "import torch; print('ABI TRUE' if torch.compiled_with_cxx11_abi() else 'ABI FALSE')"之后去flash-attn的GitHub Release页面,下载一个文件名完全匹配你环境的.whl文件吧。对我现在这个环境来说是uv pip install /path/to/flash_attn-2.8.3.post1+cu12torch2.7cxx11abiTRUE-cp311-cp311-linux_x86_64.whl。
大部分操作可以用 uv run 自动使用项目环境,无需手动 activate:
# 运行 Python 脚本(自动使用 uv 环境)
uv run python -m cs336_alignment.run_benchmarks --help
# 进入交互式 Python
uv run python
# 在 uv 环境里执行任意命令
uv run bash
# 运行测试
uv run pytest手动激活虚拟环境(例如安装本地 wheel、调试依赖时):
source .venv/bin/activate这个 MATH-12K 不开源确实难搞,不能体会到原汁原味的课程内容,不过好像也有社区总结。
下载方式见各小节,本地数据已就绪,存放在 data/ 下。
./hfd.sh Qwen/Qwen2.5-Math-1.5B \
--local-dir /root/gpufree-share/models/Qwen2.5-Math-1.5B \
-x 8 -j 6- 来源: openai/gsm8k — OpenAI 发布的 8,500 道小学数学应用题,每道需 2-8 步推理。
- 许可: MIT
- 格式: Parquet(
test-*.parquet/train-*.parquet) - 数量: train 7,473 / test 1,319(每个 variant 一样)
- 数据字段:
question(数学题文本) /answer(逐步推理 + 最终答案)
两个子集:
| 子集 | 目录 | Answer 格式特点 |
|---|---|---|
main |
data/gsm8k/main/ |
标准 CoT,推理 + <<expr=result>> 计算标注 + #### 42 最终答案 |
socratic |
data/gsm8k/socratic/ |
苏格拉底自问自答式,如 How many X? ** ... \n#### 42 |
示例(main):
Q: Natalia sold clips to 48 of her friends in April, and then she sold
half as many clips in May. How many clips did Natalia sell altogether?
A: Natalia sold 48/2 = <<48/2=24>>24 clips in May.
Natalia sold 48+24 = <<48+24=72>>72 clips altogether.
#### 72
./hfd.sh openai/gsm8k \
--dataset \
--local-dir data/gsm8kSFT 格式转换(scripts/prepare_sft_data.py):
将原始 parquet 转为 JSON/JSONL,每行包含:
question: 数学题原文answer: GSM8K 原始答案(CoT +#### 42)prompt: r1_zero 模板格式化后的 prompt(<think>前缀)response:<think> 推理过程 </think> <answer> 答案 </answer>格式
uv run python scripts/prepare_sft_data.py输出:
| 文件 | 记录数 | 用途 |
|---|---|---|
data/gsm8k/train.json / .jsonl |
7,473 | SFT 训练 |
data/gsm8k/test.json / .jsonl |
1,319 | SFT 验证 / 评估 |
- 来源: allenai/tulu-3-sft-personas-math — Allen AI 发布的合成数学指令数据集,通过 persona 增强生成 149,960 条复杂数学应用题。
- 许可: ODC-BY(研究 / 教育用途)
- 格式: Parquet(
train-00000-of-00002.parquet+train-00001-of-00002.parquet) - 数量: 149,960(仅 train,无标准 test split)
- 数据字段:
id(唯一标识) /prompt(数学题) /messages([{"role":"user","content":...}, {"role":"assistant","content":...}]) - 生成模型: GPT-4o、Claude 3.5 Sonnet
- 用途: SFT 阶段增强模型的复杂数学推理能力。与 GSM8K 相比,题目更复杂、场景更多样化(融合了各类 persona 背景)。
SFT 格式转换(scripts/prepare_sft_data.py):
uv run python scripts/prepare_sft_data.py将 messages 字段拆为 prompt(user)+ response(assistant),输出:
| 文件 | 记录数 | 用途 |
|---|---|---|
data/tulu-3-sft-personas-math/train.json / .jsonl |
149,960 | SFT 训练 |
- 来源: cais/mmlu — Massive Multitask Language Understanding,57 个学科的多选题基准。
- 许可: 见 repository(非商业用途)
- 格式: HuggingFace datasets(每科目一个子目录)
- 数量: 约 14,042 test / 1,531 val / 285 dev / 99,842 auxiliary_train
- 数据字段:
question(题目) /choices(4 个选项列表) /answer(正确答案字母 A-D)
覆盖 57 个科目,包括 elementary_mathematics、us_history、computer_science、law 等,评价模型跨领域知识。
./hfd.sh cais/mmlu \
--dataset \
--local-dir data/mmlu_hf- 来源: Anthropic/hh-rlhf — 人类偏好数据,用于训练 Reward Model / RLHF。
- 许可: MIT
- 格式: JSONL gzip(每行
{"chosen": "...", "rejected": "..."}) - 内容: 对话数据,包含 5 个子集:
| 子集 | 目录 | 说明 |
|---|---|---|
helpful-base |
data/hh/helpful-base/ |
基于 base model 的有用性偏好 |
helpful-online |
data/hh/helpful-online/ |
迭代在线 RLHF 采样数据 |
helpful-rejection-sampled |
data/hh/helpful-rejection-sampled/ |
Rejection sampling 数据 |
harmless-base |
data/hh/harmless-base/ |
无害性偏好数据 |
red-team-attempts |
data/hh/red-team-attempts/ |
红队攻击对话记录(含评分、标签) |
./hfd.sh Anthropic/hh-rlhf \
--dataset \
--local-dir data/hh- 来源: tatsu-lab/alpaca_eval — 自动化 LLM 评估基准,基于 AlpacaFarm 的 805 条指令。
- 许可: CC-BY-NC-4.0
- 格式: JSON(
alpaca_eval.json,805 条) - 数据字段:
instruction(指令) /output(参考输出) /generator(生成器标识) /dataset(来源) - 用途: 评估模型对开放性指令的回复质量(通常用 GPT-4 / 自动 judge 打分)
./hfd.sh tatsu-lab/alpaca_eval \
--dataset \
--local-dir data/alpaca_eval_hf- 来源: Bertievidgen/SimpleSafetyTests — 100 条关键安全风险测试用例。
- 许可: CC-BY-2.0
- 格式: CSV(
sst_test_cases.csv,100 条 prompt) - 危害类别: 自杀/自残、人身伤害、非法/管制物品、诈骗、儿童虐待
- 用途: 快速评估模型是否拒绝有害请求。正常模型应对全部 100 条 prompt 都拒绝回答。
./hfd.sh Bertievidgen/SimpleSafetyTests \
--dataset \
--local-dir data/simple_safety_tests使用 run_benchmarks.py 对模型进行 GSM8K 测试集评估。支持两种推理后端:vllm(默认,推荐)和 hf(fallback)。
# 安装编译依赖(仅首次需要)
apt-get update && apt-get install -y build-essential python3-dev cuda-cudart-dev-12-0
CUDA_VISIBLE_DEVICES=0 VLLM_WORKER_MULTIPROC_METHOD=spawn \
uv run python -m cs336_alignment.run_benchmarks \
--model_id /root/gpufree-share/models/Qwen2.5-Math-1.5B \
--engine vllm \
--benchmarks gsm8k \
--gsm8k_path data/gsm8k/main/test-00000-of-00001.parquet \
--output_dir outputs/baseline_qwen_math_vllm \
--max_new_tokens 512 \
--gpu_memory_utilization 0.90 \
--max_model_len 2048
VLLM_WORKER_MULTIPROC_METHOD=spawn是 vLLM V1 引擎在 Linux 下的必要条件——Python 默认fork会复制父进程 CUDA 上下文,导致子进程RuntimeError: Cannot re-initialize CUDA in forked subprocess。
uv run python -m cs336_alignment.run_benchmarks \
--model_id /root/gpufree-share/models/Qwen2.5-Math-1.5B \
--engine hf \
--benchmarks gsm8k \
--gsm8k_path data/gsm8k/main \
--output_dir outputs/baseline_qwen_math_hf \
--device cuda:0 \
--hf_batch_size 8 \
--attn_implementation eager \
--max_new_tokens 512参数说明:
--model_id— 模型路径或 HuggingFace ID--benchmarks— 评估基准。支持gsm8k、math,或用逗号组合gsm8k,math--gsm8k_path— 数据文件或目录。接受data/gsm8k、data/gsm8k/main或具体的 parquet/jsonl 文件--output_dir— 输出目录(生成summary.json和gsm8k_predictions.jsonl)--limit N— 仅跑前 N 条做快速验证--seed— 随机种子(默认 0)
输出示例:
GSM8K summary:
benchmark: gsm8k
split: test
num_examples: 1319
correct: 372
accuracy: 0.2820
parsed: 1310
parsed_ratio: 0.9932
MATH 评估在 GSM8K 的基础上多了一个关键差异:答案判分方式。GSM8K 的答案永远是最后一个数字(#### 42),但 MATH 的答案用 LaTeX 表示(\dfrac{1}{9}、\boxed{420}),需要符号级等价性判断。
评估使用 reasoning/rewards.py 中的 grade() 函数,它先用字符串归一化做快速比较,再 fallback 到 sympy 化简做数学等价性判断,可处理 \frac{1}{9} ≡ 1/9、\dfrac{1}{9} ≡ \frac{1}{9} 等场景。
# vLLM 后端(推荐)
CUDA_VISIBLE_DEVICES=0 VLLM_WORKER_MULTIPROC_METHOD=spawn \
uv run python -m cs336_alignment.run_benchmarks \
--model_id /root/gpufree-share/models/Qwen2.5-Math-1.5B \
--engine vllm \
--benchmarks math \
--math_path /root/gpufree-share/data/MATH/validation.jsonl \
--output_dir outputs/baseline_math \
--max_new_tokens 1024 \
--gpu_memory_utilization 0.90 \
--max_model_len 4096
# HF 后端(无 vLLM 时 fallback)
uv run python -m cs336_alignment.run_benchmarks \
--model_id /root/gpufree-share/models/Qwen2.5-Math-1.5B \
--engine hf \
--benchmarks math \
--math_path /root/gpufree-share/data/MATH/validation.jsonl \
--output_dir outputs/baseline_math_hf \
--device cuda:0 \
--hf_batch_size 8 \
--max_new_tokens 1024参数说明:
| 参数 | 默认值 | 说明 |
|---|---|---|
--math_path |
/root/gpufree-share/data/MATH/validation.jsonl |
MATH JSONL 文件路径 |
--benchmarks |
gsm8k |
改为 math 或 gsm8k,math 同时跑两个 |
输出包含 per-subject 和 per-level 分桶准确率:
MATH summary:
benchmark: math
num_examples: 5000
correct: 1850
accuracy: 0.37
format_rate: 1.0
GRPO 训练接受 --train_data(问题集,每行 {"problem": ..., "answer": ...} 或 {"question": ..., "answer": ...})和 --val_data(验证集)。
必备参数:
| 参数 | 默认值 | 说明 |
|---|---|---|
--model_id |
必填 | 模型路径或 HF ID。从 EI 或 SFT checkpoint 启动 |
--train_data |
必填 | 训练集 JSONL 路径(支持 MATH 格式 problem 或 GSM8K 格式 question) |
--val_data |
必填 | 验证集 JSONL 路径 |
--prompt_path |
必填 | Prompt 模板文件路径(如 cs336_alignment/prompts/r1_zero.prompt) |
双卡模式(cuda:0 训练 + cuda:1 vLLM 生成):
uv run python -m cs336_alignment.run_grpo \
--model_id /root/gpufree-share/models/Qwen2.5-Math-1.5B-EI-round3 \
--train_data /root/gpufree-share/data/MATH/train.jsonl \
--val_data /root/gpufree-share/data/MATH/validation.jsonl \
--prompt_path cs336_alignment/prompts/r1_zero.prompt \
--device cuda:0 --engine vllm --vllm_device cuda:1 --vllm_gpu_util 0.85 \
--group_size 8 --rollout_batch_size 256 \
--train_batch_size 32 --grad_accum_steps 8 \
--lr 3e-5 --n_grpo_steps 100 \
--loss_type grpo_clip --length_norm_type mask_normalize \
--kl_coef 0.04 \
--early_stopping_patience 3 \
--eval_every_steps 5 --save_every_steps 25 \
--eval_limit 200 \
--wandb_project cs336-grpo \
--output_dir outputs/grpo_v1单卡模式(HF generate 替代 vLLM):
uv run python -m cs336_alignment.run_grpo \
--model_id /root/gpufree-share/models/Qwen2.5-Math-1.5B-EI-round3 \
--train_data /root/gpufree-share/data/MATH/train.jsonl \
--val_data /root/gpufree-share/data/MATH/validation.jsonl \
--prompt_path cs336_alignment/prompts/r1_zero.prompt \
--device cuda:0 --engine hf \
--group_size 8 --rollout_batch_size 32 \
--train_batch_size 8 --grad_accum_steps 4 \
--lr 1e-5 --n_grpo_steps 100 \
--loss_type grpo_clip --length_norm_type mask_normalize \
--kl_coef 0.04 \
--early_stopping_patience 3 \
--eval_every_steps 5 --save_every_steps 25 \
--eval_limit 200 \
--wandb_project cs336-grpo \
--output_dir outputs/grpo_v1烟雾测试(快速验证流程):
uv run python -m cs336_alignment.run_grpo \
--model_id /root/gpufree-share/models/Qwen2.5-Math-1.5B-EI-round3 \
--train_data /root/gpufree-share/data/MATH/train.jsonl \
--val_data /root/gpufree-share/data/MATH/validation.jsonl \
--prompt_path cs336_alignment/prompts/r1_zero.prompt \
--device cuda:0 --engine hf \
--group_size 4 --rollout_batch_size 8 \
--train_batch_size 4 --grad_accum_steps 2 \
--n_grpo_steps 3 \
--lr 3e-5 --loss_type grpo_clip --length_norm_type mask_normalize \
--kl_coef 0.04 --early_stopping_patience 3 \
--train_limit 8 --eval_limit 10 \
--output_dir outputs/grpo_smoke参数说明:
| 参数 | 默认值 | 说明 |
|---|---|---|
--engine |
vllm |
生成引擎。vllm 需要双卡,hf 单卡可用 |
--group_size |
8 | 每问题生成 G 条回答 |
--rollout_batch_size |
128 | 每轮总回答数(需被 group_size 整除) |
--train_batch_size |
16 | 逻辑 batch 大小 |
--grad_accum_steps |
4 | 梯度累积步数。micro_batch_size = train_batch_size / grad_accum_steps |
--lr |
3e-5 | 学习率。1.5B 模型建议 1e-5~3e-5 |
--loss_type |
grpo_clip |
策略梯度类型:no_baseline / reinforce_with_baseline / grpo_clip |
--length_norm_type |
mask_normalize |
长度归一化方式。mask_normalize(Dr. GRPO)优于 mask_mean |
--kl_coef |
0.0 | KL 散度惩罚系数。0=禁用。建议 0.01~0.1 |
--early_stopping_patience |
3 | 连续 N 次 eval 不创新高时停止。0=禁用 |
--n_grpo_steps |
50 | 最大 GRPO 步数 |
--eval_every_steps |
10 | 每 N 步评估一次 |
--save_every_steps |
25 | 每 N 步保存一次 checkpoint |
--wandb_project |
cs336-grpo |
wandb 项目名 |
bash make_submission.sh