在本教程中,我们将使用AllenAI 的 Open Instruct框架,为一个紧凑的指令微调语言模型构建端到端的后训练流程。我们将经历三个主要的训练阶段:监督微调(Supervised Fine-Tuning,SFT)、直接偏好优化(Direct Preference Optimization,DPO)以及使用组相对策略优化(Group Relative Policy Optimization,GRPO)的基于可验证奖励的强化学习(Reinforcement Learning with Verifiable Rewards,RLVR)。同时,我们将原始的 Tulu 3 多 GPU 技术栈适配到 16 GB 的运行环境中。我们克隆 Open Instruct 仓库,选择性地加载其原生损失函数和实用函数,配置 LoRA 适配器,为每个训练阶段准备 GSM8K 数据集,并使用确定性验证器来评估生成的数学答案。在整个工作流程中,我们保留了 Open Instruct 的核心优化逻辑,同时将其分布式组件(如 vLLM、Ray 执行器、DeepSpeed 和异步推出队列)替换为适合 Colab 使用的轻量级 Hugging Face 和 PyTorch 实现。
import os, sys, subprocess, textwrap, json, math, random, re, ast, types, dataclasses, gc, contextlib
REPO_URL = "https://github.com/allenai/open-instruct.git"
REPO_DIR = "/content/open-instruct" if os.path.isdir("/content") else "./open-instruct"
PIP_PKGS = [
"peft", "accelerate",
"ray", "wandb", "beaker-py",
"langdetect==1.0.9", "immutabledict==1.2.0", "nltk",
"absl-py", "sympy", "antlr4-python3-runtime==4.11",
"tiktoken",
]
def sh(*args):
print("$", " ".join(args))
subprocess.run(args, check=False)
def setup():
sh(sys.executable, "-m", "pip", "install", "-q", *PIP_PKGS)
if not os.path.isdir(REPO_DIR):
sh("git", "clone", "--depth", "1", REPO_URL, REPO_DIR)
if REPO_DIR not in sys.path:
sys.path.insert(0, REPO_DIR)
os.environ.setdefault("WANDB_MODE", "disabled")
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
os.environ.setdefault("RAY_DISABLE_IMPORT_WARNING", "1")
setup()
import numpy as np
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader
from datasets import load_dataset, Dataset
from transformers import AutoModelForCausalLM, DataCollatorForSeq2Seq, get_cosine_schedule_with_warmup
from peft import LoraConfig, get_peft_model
DEV = "cuda" if torch.cuda.is_available() else "cpu"
try:
_bf16 = DEV == "cuda" and torch.cuda.is_bf16_supported(including_emulation=False)
except TypeError:
_bf16 = DEV == "cuda" and torch.cuda.get_device_properties(0).major >= 8
AMP_DTYPE = torch.bfloat16 if _bf16 else torch.float16
USE_SCALER = AMP_DTYPE is torch.float16
print(f"device={DEV} autocast dtype={AMP_DTYPE} gpu={torch.cuda.get_device_name(0) if DEV=='cuda' else '-'}")
def oi_load(relpath, names, ns=None):
src = open(os.path.join(REPO_DIR, relpath)).read()
tree = ast.parse(src)
found = {n.name: n for n in tree.body
if isinstance(n, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)) and n.name in names}
missing = set(names) - set(found)
if missing:
raise KeyError(f"{relpath}: could not find {missing} (upstream may have renamed them)")
ns = {} if ns is None else dict(ns)
ns.update({"torch": torch, "F": F, "np": np, "enum": __import__("enum"),
"dataclasses": dataclasses, "math": math, "os": os})
future = ast.parse("from __future__ import annotations").body
mod = ast.Module(body=future + [found[n] for n in names], type_ignores=[])
exec(compile(ast.fix_missing_locations(mod), f"<open_instruct:{relpath}>", "exec"), ns)
return {n: ns[n] for n in names}
_dpo = oi_load("open_instruct/dpo_utils.py", ["dpo_loss", "_get_batch_logps"])
_pf = oi_load("open_instruct/padding_free_collator.py", ["calculate_per_token_logps"])
_rl = oi_load("open_instruct/rl_utils.py", ["masked_mean"])
_mu = oi_load("open_instruct/model_utils.py", ["estimate_kl"])
_grpo = oi_load("open_instruct/grpo_utils.py", ["GRPOLossType", "compute_grpo_loss"],
ns={"model_utils": types.SimpleNamespace(**_mu)})
dpo_loss = _dpo["dpo_loss"]
get_batch_logps = _dpo["_get_batch_logps"]
per_token_logps_fn = _pf["calculate_per_token_logps"]
masked_mean = _rl["masked_mean"]
compute_grpo_loss = _grpo["compute_grpo_loss"]
GRPOLossType = _grpo["GRPOLossType"]
print("lifted from repo:", [f.__name__ for f in (dpo_loss, get_batch_logps, per_token_logps_fn,
masked_mean, compute_grpo_loss)])
from open_instruct.dataset_transformation import (
CHAT_TEMPLATES, TokenizerConfig,
sft_tulu_tokenize_and_truncate_v1, sft_tulu_filter_v1,
preference_tulu_tokenize_and_truncate_v1_2,
rlvr_tokenize_v1, visualize_token_role,
)
from open_instruct.ground_truth_utils import GSM8KVerifier, MathVerifier, IFEvalVerifierOld
我们安装了所需的轻量级依赖项,克隆了 Open Instruct 仓库,并配置了 Colab 环境以实现稳定运行。我们检测可用的 GPU 精度模式,并根据硬件能力选择 FP16 或 BF16 自动混合精度。我们还直接从仓库中提取了原始的 DPO、GRPO、掩码和对数概率函数,而无需导入其完整的分布式训练技术栈。
@dataclasses.dataclass
class CFG:
model: str = "Qwen/Qwen2.5-0.5B-Instruct"
max_seq_len: int = 640
seed: int = 42
n_sft: int = 192
sft_steps: int = 40
sft_micro_bs: int = 2
sft_accum: int = 4
sft_lr: float = 1e-4
n_dpo: int = 96
dpo_steps: int = 24
dpo_micro_bs: int = 1
dpo_accum: int = 4
dpo_lr: float = 5e-5
dpo_beta: float = 0.1
dpo_norm: bool = True
grpo_iters: int = 6
prompts_per_iter: int = 4
samples_per_prompt: int = 4
grpo_micro_bs: int = 1
grpo_inner_epochs: int = 2
grpo_lr: float = 2e-5
grpo_temperature: float = 1.0
grpo_max_new: int = 200
grpo_kl_beta: float = 0.02
clip_lower: float = 0.2
clip_higher: float = 0.272
kl_estimator: int = 2
adv_norm: str = "centered"
n_eval: int = 24
cfg = CFG()
random.seed(cfg.seed); np.random.seed(cfg.seed); torch.manual_seed(cfg.seed)
tc = TokenizerConfig(tokenizer_name_or_path=cfg.model, chat_template_name=None, use_fast=True)
tok = tc.tokenizer
print(f"\navailable CHAT_TEMPLATES: {list(CHAT_TEMPLATES)[:12]} ... ({len(CHAT_TEMPLATES)} total)")
print(f"pad={tok.pad_token!r}({tok.pad_token_id}) eos={tok.eos_token!r}({tok.eos_token_id})")
_demo = {"messages": [
{"role": "user", "content": "What is 12 * 3?"},
{"role": "assistant", "content": "12 * 3 = 36. The answer is 36."},
{"role": "user", "content": "And minus 6?"},
{"role": "assistant", "content": "36 - 6 = 30. The answer is 30."},
]}
_enc = sft_tulu_tokenize_and_truncate_v1(dict(_demo), tok, cfg.max_seq_len)
print("\n[SFT label masking — colour 0 = masked out of the loss, colour 1 = trained on]")
visualize_token_role(_enc["input_ids"].tolist(), (_enc["labels"] != -100).long().tolist(), tok)
print(f"trainable tokens: {(_enc['labels'] != -100).sum().item()}/{_enc['labels'].numel()}")
我们定义了一个集中式配置类,用于控制模型、数据集大小、学习率、批次设置以及每个训练阶段的优化参数。我们初始化 Open Instruct 分词器(Tokenizer),同时保留模型的聊天模板,并确保填充(Padding)和序列结束(End-of-Sequence)标记保持正确分离。然后,我们对一个示例对话进行分词,并可视化哪些助手生成的标记会贡献给监督训练损失。
gsm = load_dataset("openai/gsm8k", "main")
SYS = "You are a careful math assistant. Reason step by step, then finish with 'The answer is N.'"
def gsm_answer(a):
return a.split("####")[-1].strip().replace(",", "")
def gsm_solution(a):
body = a.split("####")[0].strip()
body = re.sub(r"<<.*?>>", "", body)
return f"{body}\nThe answer is {gsm_answer(a)}."
def as_messages(row):
return [{"role": "system", "content": SYS},
{"role": "user", "content": row["question"]},
{"role": "assistant", "content": gsm_solution(row["answer"])}]
train_rows = [gsm["train"][i] for i in range(cfg.n_sft + cfg.n_dpo)]
eval_rows = [gsm["test"][i] for i in range(cfg.n_eval)]
def to_lists(row):
for k in ("input_ids", "labels", "attention_mask"):
row[k] = row[k].tolist()
return row
sft_ds = Dataset.from_list([{"messages": as_messages(r)} for r in train_rows[: cfg.n_sft]])
sft_ds = sft_ds.map(lambda r: to_lists(sft_tulu_tokenize_and_truncate_v1(r, tok, cfg.max_seq_len)),
remove_columns=["messages"], desc="sft tokenize")
sft_ds = sft_ds.filter(sft_tulu_filter_v1, fn_kwargs={"tokenizer": tok}, desc="drop all-masked")
def make_pair(r):
gold = gsm_answer(r["answer"])
bad = (str(int(float(gold)) + random.choice([-10, -3, -1, 1, 2, 7]))
if gold.replace('.', '', 1).lstrip('-').isdigit() else gold + "0")
prompt = [{"role": "system", "content": SYS}, {"role": "user", "content": r["question"]}]
good_txt = gsm_solution(r["answer"])
bad_txt = good_txt.rsplit("The answer is", 1)[0] + f"The answer is {bad}."
return {"chosen": prompt + [{"role": "assistant", "content": good_txt}],
"rejected": prompt + [{"role": "assistant", "content": bad_txt}]}
dpo_ds = Dataset.from_list([make_pair(r) for r in train_rows[cfg.n_sft:]])
dpo_ds = dpo_ds.map(
lambda r: {k: (v.tolist() if torch.is_tensor(v) else v) for k, v in
preference_tulu_tokenize_and_truncate_v1_2(r, tok, cfg.max_seq_len).items()},
remove_columns=["chosen", "rejected"], desc="dpo tokenize")
rlvr_rows = [{"messages": as_messages(r)[:2], "ground_truth": gsm_answer(r["answer"]), "dataset": "gsm8k"}
for r in train_rows[: cfg.n_sft]]
rlvr_ds = Dataset.from_list(rlvr_rows).map(lambda r: rlvr_tokenize_v1(r, tok),
remove_columns=["messages"], desc="rlvr tokenize")
我们加载了 GSM8K 数据集,并格式化了系统消息,明确指示模型逐步推理并以