AllenAI、Open Instruct で Tulu 3 のポストトレーニングパイプラインを公開
本文の状態
日本語全文を表示中
詳細モードで約12分の本文を読めます。
同じ出来事の情報源
この情報源を基点に整理
MarkTechPost
AllenAI は Open Instruct フレームワークを用いた Tulu 3 の後学習パイプラインを構築し、16GB の環境でも SFT、DPO、GRPO を実行可能な軽量版を提供した。
Continue in AI NEW LAB
このニュースを、実務の判断につなげる
AI NEW LABで、試したことや先に確認したい条件を共有できます。まずはログインなしで読めます。
AI NEW LABで論点を見るAI深層分析を開く2026年8月13日 03:21
AI深層分析
キーポイント
リソース制約下でのトレーニング実装
元の Tulu 3 スタックを 16GB のランタイム環境に適合させるため、vLLM や Ray などの分散コンポーネントを軽量な Hugging Face および PyTorch 実装へ置き換えた。
多段階学習パイプラインの構築
教師あり微調整(SFT)、直接選好最適化(DPO)、検証可能な報酬を用いた強化学習(GRPO)という 3 つの主要なトレーニングステージを順次実行するエンドツーエンドのワークフローを提供した。
数学的推論能力の評価手法
GSM8K データセットを用いて各学習ステージを準備し、生成された数式回答に対して決定論的な検証器(deterministic verifiers)を使用して評価を行う仕組みを組み込んだ。
Colab 環境での利用可能性
非同期ロールアウトキューや分散処理の代替として軽量なライブラリを採用することで、Google Colab などの制約のあるクラウド環境でも本格的な後学習パイプラインを実行可能にした。
依存パッケージとリポジトリのセットアップ
peft, accelerate, ray, wandbなどの主要ライブラリをインストールし、AllenAIのopen-instructリポジトリをクローンして環境変数を設定する。
重要な引用
We move through three major training stages: Supervised Fine-Tuning, Direct Preference Optimization, and Reinforcement Learning with Verifiable Rewards using GRPO
adapting the original multi-GPU Tulu 3 stack to fit within a 16 GB runtime
replacing distributed components such as vLLM, Ray actors, DeepSpeed, and asynchronous rollout queues with lightweight Hugging Face and PyTorch implementations
REPO_URL = "https://github.com/allenai/open-instruct.git"
編集コメントを表示
編集コメント
本記事は、大規模な計算資源を必要とする最新の RL 学習手法を、個人が利用可能な環境で実行するための具体的な実装コードを提供している。技術的な詳細への深い理解が求められるが、実践的な学習パイプラインの構築を目指す開発者にとって極めて有用なリソースである。
Source Article
元記事を日本語で読む
本文に関係しない購読案内、埋め込み通知、サイト内プロモーションは除いています。
このチュートリアルでは、AllenAI の Open Instruct フレームワークを活用し、コンパクトな指令微調整済み言語モデル向けのエンドツーエンドのポストトレーニングパイプラインを構築します。主な学習プロセスは 3 つの段階に分かれています。
まず「教師あり微調整(Supervised Fine-Tuning)」、次に「直接選好最適化(Direct Preference Optimization)」、そして GRPO を用いた検証可能な報酬による強化学習です。
これらの手順を通じて、オリジナルのマルチ GPU 構成である Tulu 3 スタックを 16 GB のメモリ環境で動作するように調整します。具体的には Open Instruct リポジトリをクローンし、必要なネイティブ損失関数やユーティリティ関数を選択的に読み込みます。また、LoRA アダプタを設定し、各学習段階に合わせた GSM8K データの準備を行います。
生成された数学的解答の評価には、決定論的な検証器(verifier)を使用します。このワークフロー全体を通じて、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"", "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)})
必要な軽量な依存関係をインストールし、Open Instruct リポジトリをクローンして、Colab 環境を安定実行のために設定します。利用可能な GPU の精度モードを検出し、ハードウェアの能力に応じて FP16 または BF16 の自動キャストを選択します。また、フル分散トレーニングスタックをインポートするのではなく、DPO、GRPO、マスキング、対数確率計算の元の関数をリポジトリから直接抽出して使用します。
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@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) を実行し、以下の出力を確認します。
[SFT label masking — colour 0 = masked out of the loss, colour 1 = trained on]トークン化された入力 ID とラベル(-100 以外が訓練対象)を可視化する関数 visualize_token_role を呼び出し、どのトークンが教師あり学習の損失計算に寄与しているかを確認します。さらに、訓練可能なトークンの数を _enc['labels'] != -100 の条件でカウントし、全体のトークン数に対する割合を表示します。
各トレーニングステージにおいて、モデル設定やデータセットサイズ、学習率、バッチ設定、最適化パラメータなどを一元管理する設定クラスを定義しています。Open Instruct トークナイザーを初期化する際は、モデルのチャットテンプレートを引き継ぎつつ、パディングトークンと終端トークンが正しく区別されるよう注意して処理を行います。その後、サンプル会話データをトークン化し、アシスタント側のトークンが教師あり学習の損失計算にどのように関与しているかを可視化します。
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")
print(f"\nsft={len(sft_ds)} dpo={len(dpo_ds)} rlvr={len(rlvr_ds)}")
VERIFIERS = {"gsm8k": GSM8KVerifier(), "math": MathVerifier(), "ifeval_old": IFEvalVerifierOld()}
print("\n[verifier smoke test]")
print(" gsm8k :", VERIFIERS"gsm8k".score)
print(" gsm8k :", VERIFIERS"gsm8k".score)
print(" math :", VERIFIERS"math".score)
print(" ifeval:", VERIFIERS["ifeval_old"]([], "one two three four five six seven",
json.dumps({"func_name": "validate_word_constraint",</article>}
「N」を 6、「quantifier」を「少なくとも」と設定したスコア計算後、検証バッチ処理関数を実行します。この関数は、生成された回答、正解、およびソース情報を引数として受け取り、各質問に対して適切な検証器(GSM8K データセット用またはソース固有の検証器)を選択してスコアを算出します。最終的に、重み付けされたスコアの配列を返却します。
SFT(教師あり学習)、DPO(直接最適化)、RLVR(検証付き強化学習)の各トレーニングにおいて一貫した対話形式で利用できるよう、GSM8K データセットの質問と解答を再構成しました。具体的には、教師あり学習用の例、意図的に最終回答を誤った比較ペア、そして構造化された正解ラベルを持つ検証器用プロンプトを作成しています。さらに、Open Instruct の GSM8K 専用、数学問題専用、指示従順性評価専用の各検証器を初期化し、生成された回答に対して決定論的なスコアリングを実施します。
model = AutoModelForCausalLM.from_pretrained(cfg.model, dtype=torch.float32).to(DEV)
model.config.use_cache = False
if len(tok) > model.get_input_embeddings().weight.shape[0]:
model.resize_token_embeddings(len(tok))
def _patch_peft_torchao():
import importlib
for mod in ("peft.import_utils", "peft.tuners.lora.torchao",
"peft.tuners.lora.model", "peft.tuners.lora.layer"):
try:
m = importlib.import_module(mod)
except Exception:
continue
if hasattr(m, "is_torchao_available"):
m.is_torchao_available = lambda: False
_patch_peft_torchao()
model = get_peft_model(model, LoraConfig(
r=32, lora_alpha=64, lora_dropout=0.05, bias="none", task_type="CAUSAL_LM",
model.print_trainable_parameters()
TRAINABLE = [p for p in model.parameters() if p.requires_grad]
@contextlib.contextmanager
def with_cache():
old = model.config.use_cache
model.config.use_cache = True
try:
yield
finally:
model.config.use_cache = old
def amp():
return torch.autocast(device_type="cuda", dtype=AMP_DTYPE) if DEV == "cuda" \
else torch.autocast(device_type="cpu", enabled=False)
def new_opt(lr, steps):
opt = torch.optim.AdamW(TRAINABLE, lr=lr, weight_decay=0.0, betas=(0.9, 0.999))
sched = get_cosine_schedule_with_warmup(opt, int(0.05 * steps) + 1, steps)
scaler = torch.amp.GradScaler("cuda", enabled=USE_SCALER)
return opt, sched, scaler
def step_opt(opt, sched, scaler):
scaler.unscale_(opt)torch.nn.utils.clip_grad_norm_(TRAINABLE, 1.0)
scaler.step(opt); scaler.update(); sched.step(); opt.zero_grad(set_to_none=True)
@torch.no_grad()
def evaluate(tag, rows, max_new=256):
model.eval()
tok.padding_side = "left"
correct, bs = 0.0, 4
for i in range(0, len(rows), bs):
chunk = rows[i:i + bs]
prompts = [tok.apply_chat_template(
[{"role": "system", "content": SYS}, {"role": "user", "content": r["question"]}],
add_generation_prompt=True, tokenize=False) for r in chunk]
enc = tok(prompts, return_tensors="pt", padding=True, add_special_tokens=False).to(DEV)
with amp(), with_cache():
out = model.generate(**enc, max_new_tokens=max_new, do_sample=False,
pad_token_id=tok.pad_token_id)
texts = tok.batch_decode(out[:, enc["input_ids"].shape[1]:], skip_special_tokens=True)
correct += verify_batch(texts, [gsm_answer(r["answer"]) for r in chunk],
["gsm8k"] * len(chunk)).sum()
acc = correct / len(rows)
print(f" [eval:{tag}] verifier accuracy = {acc:.3f} ({int(correct)}/{len(rows)})")
model.train(); tok.padding_side = "right"
return acc
print("\n" + "=" * 90); print("BASELINE"); print("=" * 90)
base_acc = evaluate("base", eval_rows)
Qwen の指令モデルを読み込み、アテンション層とフィードフォワード投影層に LoRA アダプターを適用して、最適化対象を学習可能なアダプターパラメータに限定します。混合精度実行、勾配スケーリング、勾配クリッピング、学習率スケジューリング、生成時の一時的 KV キャッシュ活性化を設定します。その後、未訓練のベースラインモデルで GSM8K をグリディックデコーディングと検証器による回答精度評価を用いてテストします。
"\n" + "=" * 90); print("STAGE 1 — SFT"); print("=" * 90)
sft_collate = DataCollatorForSeq2Seq(tokenizer=tok, padding="longest", label_pad_token_id=-100)
sft_dl = DataLoader(sft_ds, batch_size=cfg.sft_micro_bs, shuffle=True, collate_fn=sft_collate, drop_last=True)
opt, sched, scaler = new_opt(cfg.sft_lr, cfg.sft_steps)
model.train(); it, step, run = iter(sft_dl), 0, 0.0
while step < cfg.sft_steps:
for batch in it:
outputs = model(**batch, return_dict=True)
loss = outputs.loss
scaler.scale(loss).backward()
if (step + 1) % cfg.gradient_accumulation == 0:
scaler.unscale_(opt)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=cfg.max_grad_norm)
scaler.step(opt)
scaler.update()
sched.step()
opt.zero_grad(set_to_none=True)
step += 1
run += loss.item()
if step % cfg.log_steps == 0:
print(f"{step}/{cfg.sft_steps} loss {run/cfg.log_steps:.4f} lr {sched.get_last_lr()[0]:.2e}")
run = 0.0
sft_acc = evaluate("after-sft", eval_rows)
パディングされた SFT の DataLoader を構築し、トークン化された GSM8K の会話データに対して勾配累積を用いて LoRA アダプターを学習します。最適化には、マスクされていないアシスタント応答トークンのみに対して計算される交差エントロピー損失を使用します。この段階全体を通じてトレーニング損失と学習率を追跡し、教師あり微調整後のモデルを更新して評価を行います。
print("\n" + "=" * 90); print("STAGE 2 — DPO (dpo_norm)"); print("=" * 90)
def pad_side(seqs, pad, maxlen):
return torch.tensor([s + [pad] * (maxlen - len(s)) for s in seqs], dtype=torch.long)
def dpo_collate(features):
out = {}
for pfx in ("chosen", "rejected"):
L = max(len(f[f"{pfx}_input_ids"]) for f in features)
out[f"{pfx}_input_ids"] = pad_side([f[f"{pfx}_input_ids"] for f in features], tok.pad_token_id, L)
out[f"{pfx}_labels"] = pad_side([f[f"{pfx}_labels"] for f in features], -100, L)
out[f"{pfx}_attention_mask"] = pad_side([f[f"{pfx}_attention_mask"] for f in features], 0, L)
return out
def seq_logps(input_ids, attn, labels):
with amp():
logits = model(input_ids=input_ids, attention_mask=attn).logits
ptl = per_token_logps_fn(logits, labels)
return get_batch_logps(ptl, labels, average_log_prob=cfg.dpo_norm)
dpo_dl = DataLoader(dpo_ds, batch_size=cfg.dpo_micro_bs, shuffle=True, collate_fn=dpo_collate, drop_last=True)
opt, sched, scaler = new_opt(cfg.dpo_lr, cfg.dpo_steps)
it, step = iter(dpo_dl), 0
while step < cfg.dpo_steps:
for batch in it:
with amp():
r_c = seq_logps(batch["chosen_input_ids"], batch["chosen_attention_mask"], batch["chosen_labels"])
r_r = seq_logps(batch["rejected_input_ids"], batch["rejected_attention_mask"], batch["rejected_labels"])
loss = -torch.nn.functional.logsigmoid(r_c - r_r).mean()
agg["loss"] += loss.item() / cfg.dpo_accum
agg["acc"] += (r_c > r_r).float().mean().item() / cfg.dpo_accum
agg["margin"] += (r_c - r_r).mean().item() / cfg.dpo_accum
step_opt(opt, sched, scaler); step += 1
if step % 8 == 0 or step == 1:
print(f" dpo step {step:>3}/{cfg.dpo_steps} loss {agg['loss']:.4f} "
f"reward_acc {agg['acc']:.2f} margin {agg['margin']:+.3f}")
dpo_acc = evaluate("after-dpo", eval_rows)
選ばれた回答と拒否された回答は別々にバッチ処理し、Open Instruct のネイティブユーティリティを用いて長さ正規化されたシーケンス対数尤度を計算します。アクティブな LoRA ポリシーを凍結されたベース参照ポリシーと比較し、レポジトリが提供する DPO 損失関数でモデルを最適化します。DPO 後の検証器性能を測定する前に、選好精度、報酬マージン、トレーニング損失を追跡します。
print("\n" + "=" * 90); print("STAGE 3 — RLVR / GRPO"); print("=" * 90)
grpo_cfg = types.SimpleNamespace(loss_fn=GRPOLossType.dapo, clip_lower=cfg.clip_lower,
clip_higher=cfg.clip_higher, kl_estimator=cfg.kl_estimator)
_gen_eos = getattr(getattr(model, "generation_config", None), "eos_token_id", None)
_terms = {tok.eos_token_id, tok.pad_token_id}
_terms |= set(_gen_eos) if isinstance(_gen_eos, (list, tuple)) else {_gen_eos}
TERMINATORS = torch.tensor(sorted(t for t in _terms if t is not None), device=DEV)
def token_logps(seq, attn, temperature, grad=True):
pos = (attn.cumsum(-1) - 1).clamp(min=0)
ctx = torch.enable_grad() if grad else torch.no_grad()
with ctx, amp():
logits = model(input_ids=seq, attention_mask=attn, position_ids=pos).logits
return per_token_logps_fn(logits / temperature, seq)
def rollout(batch_rows):
G = cfg.samples_per_prompt
ids = [r["input_ids_prompt"] for r in batch_rows]
P = max(len(x) for x in ids)
pin = torch.tensor([[tok.pad_
原文を表示
In this tutorial, we build an end-to-end post-training pipeline for a compact instruction-tuned language model using AllenAI’s Open Instruct framework. We move through three major training stages: Supervised Fine-Tuning, Direct Preference Optimization, and Reinforcement Learning with Verifiable Rewards using GRPO, while adapting the original multi-GPU Tulu 3 stack to fit within a 16 GB runtime. We clone the Open Instruct repository, selectively load its native loss and utility functions, configure LoRA adapters, prepare GSM8K data for each training stage, and use deterministic verifiers to evaluate generated mathematical answers. Throughout the workflow, we preserve the core optimization logic of Open Instruct while replacing distributed components such as vLLM, Ray actors, DeepSpeed, and asynchronous rollout queues with lightweight Hugging Face and PyTorch implementations suitable for Colab.
Copy CodeCopiedUse a different Browser
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
We install the required lightweight dependencies, clone the Open Instruct repository, and configure the Colab environment for stable execution. We detect the available GPU precision mode and select either FP16 or BF16 autocasting based on the hardware capabilities. We also extract the original DPO, GRPO, masking, and log-probability functions directly from the repository without importing its full distributed training stack.
Copy CodeCopiedUse a different Browser
@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()}")
We define a centralized configuration class that controls the model, dataset sizes, learning rates, batch settings, and optimization parameters for every training stage. We initialize the Open Instruct tokenizer while preserving the model’s chat template and ensuring that padding and end-of-sequence tokens remain correctly separated. We then tokenize a sample conversation and visualize which assistant tokens contribute to the supervised training loss.
Copy CodeCopiedUse a different Browser
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")
print(f"\nsft={len(sft_ds)} dpo={len(dpo_ds)} rlvr={len(rlvr_ds)}")
VERIFIERS = {"gsm8k": GSM8KVerifier(), "math": MathVerifier(), "ifeval_old": IFEvalVerifierOld()}
print("\n[verifier smoke test]")
print(" gsm8k :", VERIFIERS"gsm8k".score)
print(" gsm8k :", VERIFIERS"gsm8k".score)
print(" math :", VERIFIERS"math".score)
print(" ifeval:", VERIFIERS["ifeval_old"]([], "one two three four five six seven",
json.dumps({"func_name": "validate_word_constraint",
"N": 6, "quantifier": "at least"})).score)
def verify_batch(responses, ground_truths, sources, tokenized=None):
out = []
for i, (resp, gt, src) in enumerate(zip(responses, ground_truths, sources)):
v = VERIFIERS.get(src, VERIFIERS["gsm8k"])
out.append(v(tokenized[i] if tokenized else [], resp, gt).score * v.weight)
return np.array(out, dtype=np.float32)
We load GSM8K and transform its questions and solutions into a consistent conversational format for SFT, DPO, and RLVR training. We create supervised examples, preference pairs with deliberately incorrect final answers, and verifier-ready prompts with structured ground-truth labels. We also initialize Open Instruct’s GSM8K, mathematical, and instruction-following verifiers and use them to score generated responses deterministically.
Copy CodeCopiedUse a different Browser
model = AutoModelForCausalLM.from_pretrained(cfg.model, dtype=torch.float32).to(DEV)
model.config.use_cache = False
if len(tok) > model.get_input_embeddings().weight.shape[0]:
model.resize_token_embeddings(len(tok))
def _patch_peft_torchao():
import importlib
for mod in ("peft.import_utils", "peft.tuners.lora.torchao",
"peft.tuners.lora.model", "peft.tuners.lora.layer"):
try:
m = importlib.import_module(mod)
except Exception:
continue
if hasattr(m, "is_torchao_available"):
m.is_torchao_available = lambda: False
_patch_peft_torchao()
model = get_peft_model(model, LoraConfig(
r=32, lora_alpha=64, lora_dropout=0.05, bias="none", task_type="CAUSAL_LM",
model.print_trainable_parameters()
TRAINABLE = [p for p in model.parameters() if p.requires_grad]
@contextlib.contextmanager
def with_cache():
old = model.config.use_cache
model.config.use_cache = True
try:
yield
finally:
model.config.use_cache = old
def amp():
return torch.autocast(device_type="cuda", dtype=AMP_DTYPE) if DEV == "cuda" \
else torch.autocast(device_type="cpu", enabled=False)
def new_opt(lr, steps):
opt = torch.optim.AdamW(TRAINABLE, lr=lr, weight_decay=0.0, betas=(0.9, 0.999))
sched = get_cosine_schedule_with_warmup(opt, int(0.05 * steps) + 1, steps)
scaler = torch.amp.GradScaler("cuda", enabled=USE_SCALER)
return opt, sched, scaler
def step_opt(opt, sched, scaler):
scaler.unscale_(opt)
torch.nn.utils.clip_grad_norm_(TRAINABLE, 1.0)
scaler.step(opt); scaler.update(); sched.step(); opt.zero_grad(set_to_none=True)
@torch.no_grad()
def evaluate(tag, rows, max_new=256):
model.eval()
tok.padding_side = "left"
correct, bs = 0.0, 4
for i in range(0, len(rows), bs):
chunk = rows[i:i + bs]
prompts = [tok.apply_chat_template(
[{"role": "system", "content": SYS}, {"role": "user", "content": r["question"]}],
add_generation_prompt=True, tokenize=False) for r in chunk]
enc = tok(prompts, return_tensors="pt", padding=True, add_special_tokens=False).to(DEV)
with amp(), with_cache():
out = model.generate(**enc, max_new_tokens=max_new, do_sample=False,
pad_token_id=tok.pad_token_id)
texts = tok.batch_decode(out[:, enc["input_ids"].shape[1]:], skip_special_tokens=True)
correct += verify_batch(texts, [gsm_answer(r["answer"]) for r in chunk],
["gsm8k"] * len(chunk)).sum()
acc = correct / len(rows)
print(f" [eval:{tag}] verifier accuracy = {acc:.3f} ({int(correct)}/{len(rows)})")
model.train(); tok.padding_side = "right"
return acc
print("\n" + "=" * 90); print("BASELINE"); print("=" * 90)
base_acc = evaluate("base", eval_rows)
We load the Qwen instruction model, apply LoRA adapters to its attention and feed-forward projection layers, and restrict optimization to the trainable adapter parameters. We configure mixed-precision execution, gradient scaling, gradient clipping, learning-rate scheduling, and temporary KV-cache activation for generation. We then evaluate the untrained baseline on GSM8K using greedy decoding and verifier-based answer accuracy.
Copy CodeCopiedUse a different Browser
print("\n" + "=" * 90); print("STAGE 1 — SFT"); print("=" * 90)
sft_collate = DataCollatorForSeq2Seq(tokenizer=tok, padding="longest", label_pad_token_id=-100)
sft_dl = DataLoader(sft_ds, batch_size=cfg.sft_micro_bs, shuffle=True, collate_fn=sft_collate, drop_last=True)
opt, sched, scaler = new_opt(cfg.sft_lr, cfg.sft_steps)
model.train(); it, step, run = iter(sft_dl), 0, 0.0
while step < cfg.sft_steps:
for _ in range(cfg.sft_accum):
try:
batch = next(it)
except StopIteration:
it = iter(sft_dl); batch = next(it)
batch = {k: v.to(DEV) for k, v in batch.items()}
with amp():
loss = model(**batch).loss / cfg.sft_accum
scaler.scale(loss).backward()
run += loss.item()
step_opt(opt, sched, scaler); step += 1
if step % 10 == 0 or step == 1:
print(f" sft step {step:>3}/{cfg.sft_steps} loss {run:.4f} lr {sched.get_last_lr()[0]:.2e}")
run = 0.0
sft_acc = evaluate("after-sft", eval_rows)
We construct a padded SFT DataLoader and train the LoRA adapters on tokenized GSM8K conversations using gradient accumulation. We optimize the model with cross-entropy loss calculated only over the unmasked assistant response tokens. We track the training loss and learning rate throughout the stage and evaluate the updated model after supervised fine-tuning.
Copy CodeCopiedUse a different Browser
print("\n" + "=" * 90); print("STAGE 2 — DPO (dpo_norm)"); print("=" * 90)
def pad_side(seqs, pad, maxlen):
return torch.tensor([s + [pad] * (maxlen - len(s)) for s in seqs], dtype=torch.long)
def dpo_collate(features):
out = {}
for pfx in ("chosen", "rejected"):
L = max(len(f[f"{pfx}_input_ids"]) for f in features)
out[f"{pfx}_input_ids"] = pad_side([f[f"{pfx}_input_ids"] for f in features], tok.pad_token_id, L)
out[f"{pfx}_labels"] = pad_side([f[f"{pfx}_labels"] for f in features], -100, L)
out[f"{pfx}_attention_mask"] = pad_side([f[f"{pfx}_attention_mask"] for f in features], 0, L)
return out
def seq_logps(input_ids, attn, labels):
with amp():
logits = model(input_ids=input_ids, attention_mask=attn).logits
ptl = per_token_logps_fn(logits, labels)
return get_batch_logps(ptl, labels, average_log_prob=cfg.dpo_norm)
dpo_dl = DataLoader(dpo_ds, batch_size=cfg.dpo_micro_bs, shuffle=True, collate_fn=dpo_collate, drop_last=True)
opt, sched, scaler = new_opt(cfg.dpo_lr, cfg.dpo_steps)
it, step = iter(dpo_dl), 0
while step < cfg.dpo_steps:
agg = {"loss": 0.0, "acc": 0.0, "margin": 0.0}
for _ in range(cfg.dpo_accum):
try:
b = next(it)
except StopIteration:
it = iter(dpo_dl); b = next(it)
b = {k: v.to(DEV) for k, v in b.items()}
with torch.no_grad(), model.disable_adapter():
ref_c = seq_logps(b["chosen_input_ids"], b["chosen_attention_mask"], b["chosen_labels"])
ref_r = seq_logps(b["rejected_input_ids"], b["rejected_attention_mask"], b["rejected_labels"])
pol_c = seq_logps(b["chosen_input_ids"], b["chosen_attention_mask"], b["chosen_labels"])
pol_r = seq_logps(b["rejected_input_ids"], b["rejected_attention_mask"], b["rejected_labels"])
losses, r_c, r_r = dpo_loss(pol_c, pol_r, ref_c, ref_r, beta=cfg.dpo_beta, label_smoothing=0.0)
loss = losses.mean() / cfg.dpo_accum
scaler.scale(loss).backward()
agg["loss"] += loss.item()
agg["acc"] += (r_c > r_r).float().mean().item() / cfg.dpo_accum
agg["margin"] += (r_c - r_r).mean().item() / cfg.dpo_accum
step_opt(opt, sched, scaler); step += 1
if step % 8 == 0 or step == 1:
print(f" dpo step {step:>3}/{cfg.dpo_steps} loss {agg['loss']:.4f} "
f"reward_acc {agg['acc']:.2f} margin {agg['margin']:+.3f}")
dpo_acc = evaluate("after-dpo", eval_rows)
We batch the chosen and rejected responses separately and calculate their length-normalized sequence log probabilities with Open Instruct’s native utilities. We compare the active LoRA policy against the frozen base reference policy and optimize the model using the repository’s DPO loss. We monitor preference accuracy, reward margins, and training loss before measuring the model’s post-DPO verifier performance.
Copy CodeCopiedUse a different Browser
print("\n" + "=" * 90); print("STAGE 3 — RLVR / GRPO"); print("=" * 90)
grpo_cfg = types.SimpleNamespace(loss_fn=GRPOLossType.dapo, clip_lower=cfg.clip_lower,
clip_higher=cfg.clip_higher, kl_estimator=cfg.kl_estimator)
_gen_eos = getattr(getattr(model, "generation_config", None), "eos_token_id", None)
_terms = {tok.eos_token_id, tok.pad_token_id}
_terms |= set(_gen_eos) if isinstance(_gen_eos, (list, tuple)) else {_gen_eos}
TERMINATORS = torch.tensor(sorted(t for t in _terms if t is not None), device=DEV)
def token_logps(seq, attn, temperature, grad=True):
pos = (attn.cumsum(-1) - 1).clamp(min=0)
ctx = torch.enable_grad() if grad else torch.no_grad()
with ctx, amp():
logits = model(input_ids=seq, attention_mask=attn, position_ids=pos).logits
return per_token_logps_fn(logits / temperature, seq)
def rollout(batch_rows):
G = cfg.samples_per_prompt
ids = [r["input_ids_prompt"] for r in batch_rows]
P = max(len(x) for x in ids)
pin = torch.tensor([[tok.pad_
関連記事
今日のまとめ
AIデイリーブリーフで今日の重要ニュースをまとめ読み