基盤モデルのための強化学習の実装
『基盤モデルのための強化学習』第4章・第5章の数式と、TRL の実装の対応を補足します。記号と各手法の導出は本書を参照してください。
なお、TRL の参照コミットは、2026年10月2日時点で最新版の 4b4c05aとします。ただし、PPO は実装が含まれる v1.8.0 を参照します。
4.4.2節 / 式 (4.8)RLHF:報酬モデルの学習
本書の式 (4.8) の損失関数は次の通りです。
報酬関数は、RewardTrainer で学習できます。損失関数の計算は次のように記述できます。
import torch.nn.functional as F
loss = -F.logsigmoid(rewards_chosen - rewards_rejected).mean()
rewards_chosen と rewards_rejected は、それぞれ式中の \(r_\phi(x,y_w)\) と \(r_\phi(x,y_\ell)\) に対応します。
選好データセット preference_dataset は、以下のようなデータを複数件収集することによって作成します。prompt・chosen・rejected は、本書の \(x,y_w,y_\ell\) に対応します。
from datasets import Dataset
preference_dataset = Dataset.from_list([
{
"prompt": "日本の首都は?",
"chosen": "東京です。",
"rejected": "札幌です。",
},
])
SFT 済みモデルを sft_model、対応するトークナイザを tokenizer とすると、報酬モデルを選好データセットから学習することが可能です。
from trl import RewardConfig, RewardTrainer
trainer = RewardTrainer(
model=sft_model,
processing_class=tokenizer,
args=RewardConfig(
output_dir="outputs/reward-model",
),
train_dataset=preference_dataset,
)
trainer.train()
trainer.save_model("outputs/reward-model")
4.4.3節 / 式 (4.9)・(4.10)、2.3.5.2節 / 式 (2.29)RLHF:PPO の実装
学習した報酬モデルを使い、本書の式 (4.9) の目的関数を最大化します。
ここでは、TRL v1.8.0 の PPOTrainer を使います。式 (4.10) のKLペナルティ付き報酬からGAEでアドバンテージを計算し、式 (2.29) のクリッピングを適用します。方策の損失計算の要点は次の通りです。
import torch
ratio = torch.exp(new_logprobs - old_logprobs)
clipped_ratio = torch.clamp(ratio, 1 - cliprange, 1 + cliprange)
per_token_loss = -torch.minimum(ratio * advantages, clipped_ratio * advantages)
policy_loss = (per_token_loss * mask).sum() / mask.sum()
ratio は式 (2.29) の \(\pi_\theta(a_t\mid s_t)/\pi_{\theta_{\mathrm{old}}}(a_t\mid s_t)\)、advantages は \(\hat{A}_t\)、cliprange は \(\epsilon\) に対応します。
以下のコードでは、4つのモデルを用意します。
policy:SFT 済みの方策。PPO で更新します。ref_policy:SFT 済みの方策のコピー。参照方策として固定します。reward_model:前節で学習した報酬モデル。PPO の学習中は固定します。value_model:価値関数を推定するモデル。学習済みの報酬モデルreward_modelで初期化することが多いです。
PPO はアクター・クリティック法の一つです。以下のようなソースコードで、policy(アクター)と value_model(クリティック)を同時に学習します。
from trl.experimental.ppo import PPOConfig, PPOTrainer
trainer = PPOTrainer(
model=policy,
ref_model=ref_policy,
reward_model=reward_model,
value_model=value_model,
processing_class=tokenizer,
args=PPOConfig(
output_dir="outputs/ppo",
kl_coef=0.05,
kl_estimator="k1",
cliprange=0.2,
stop_token="eos",
num_sample_generations=0,
),
train_dataset=prompt_dataset,
)
trainer.train()
kl_coef が式 (4.9)・(4.10) の \(\beta\) に対応します。kl_estimator="k1" では、式 (4.10) の対数確率比をKLペナルティに使います。ref_policy と reward_model は固定し、方策と価値モデルを更新します。
4.5.1節 / 式 (4.19)DPO の実装
本書の式 (4.19) の損失関数は次の通りです。
式 (4.19) は、DPOTrainer の loss_type="sigmoid" に対応します。損失計算の要点は次の通りです。
import torch.nn.functional as F
chosen_logratios = chosen_logps - ref_chosen_logps
rejected_logratios = rejected_logps - ref_rejected_logps
delta_score = chosen_logratios - rejected_logratios
loss = -F.logsigmoid(beta * delta_score).mean()
chosen_logratios は式中の \(\log\frac{\pi_\theta(y_w\mid x)}{\pi_{\mathrm{ref}}(y_w\mid x)}\)、rejected_logratios は \(\log\frac{\pi_\theta(y_\ell\mid x)}{\pi_{\mathrm{ref}}(y_\ell\mid x)}\) に対応します。その差が delta_score なので、beta * delta_score が \(\sigma\) の引数になります。-F.logsigmoid(...).mean() で負の対数を取り、期待値をミニバッチ平均で計算します。
以下ではSFT済みモデル sft_model と対応する tokenizer を用意済みとします。beta は本書の \(\beta\) に対応します。選好データセット preference_dataset を用いて、以下のようなソースコードで学習します。
from trl import DPOConfig, DPOTrainer
trainer = DPOTrainer(
model=sft_model,
processing_class=tokenizer,
args=DPOConfig(
output_dir="outputs/dpo",
loss_type="sigmoid",
beta=0.1,
),
train_dataset=preference_dataset,
)
trainer.train()
4.5.4.2節の IPO も同じ DPOTrainer で実装できます。上の loss_type="sigmoid" を loss_type="ipo" に変更します。他の DPO の派生手法についても、loss_type を変更することで簡単に実装できます。詳細は、公式ドキュメントを参照して下さい。
4.5.4.3節 / 式 (4.29)〜(4.32)KTO の実装
KTO には KTOTrainer を使います。データは prompt・completion・label の3列とし、本書の \(\mathcal{Y}^+(x),\mathcal{Y}^-(x)\) に対応するラベルを True / False で渡します。ここで、beta は本書の \(\beta\)、desirable_weight と undesirable_weight は式 (4.32) の \(w_+,w_-\) に対応します。
from trl import KTOConfig, KTOTrainer
trainer = KTOTrainer(
model=sft_model,
processing_class=tokenizer,
args=KTOConfig(
output_dir="outputs/kto",
loss_type="kto",
beta=0.1,
desirable_weight=1.0,
undesirable_weight=1.0,
per_device_train_batch_size=4,
train_sampling_strategy="sequential",
),
train_dataset=feedback_dataset,
)
trainer.train()
5.3.1節 / 式 (5.1)・(5.4)〜(5.7)GRPO の実装
本書の式 (5.4) の目的関数は次の通りです。TRLでは、その符号を反転した損失を最小化します。
GRPOTrainer で、loss_type="grpo"、importance_sampling_level="token" を指定します。アウトカム報酬を使う場合の計算の要点は次の通りです。
import torch
grouped_rewards = rewards.view(-1, G)
advantages = (grouped_rewards - grouped_rewards.mean(-1, keepdim=True)) / (
grouped_rewards.std(-1, keepdim=True) + 1e-4
)
advantages = advantages.reshape(-1, 1)
ratio = (per_token_logps - old_per_token_logps).exp()
clipped_ratio = ratio.clamp(1 - epsilon, 1 + epsilon)
log_ref_ratio = ref_per_token_logps - per_token_logps
kl = log_ref_ratio.exp() - log_ref_ratio - 1
per_token_loss = -torch.minimum(
ratio * advantages, clipped_ratio * advantages
) + beta * kl
loss = ((per_token_loss * mask).sum(-1) / mask.sum(-1).clamp(min=1)).mean()
advantages は式 (5.1) の \(\widehat{A}_{i,t}\)、ratio は式 (5.5) の \(w_{i,t}(\theta)\)、kl は式 (5.7) の推定量です。old_per_token_logps は生成時の方策 \(\pi_{\theta_{\mathrm{old}}}\)、ref_per_token_logps は参照方策 \(\pi_{\mathrm{ref}}\) のトークン対数確率に対応します。
学習データには、prompt と ground_truth が含まれ、それぞれ本書の \(x\) と \(y^\star\) に対応します。
from datasets import Dataset
train_dataset = Dataset.from_list([
{
"prompt": "20未満の素数の合計はいくらですか? 数字だけで答えてください。",
"ground_truth": "77",
},
])
def reward_func(completions, ground_truth, **kwargs):
return [
float(completion.strip() == answer)
for completion, answer in zip(completions, ground_truth)
]
completions は生成された回答のリストです。データの追加列 ground_truth も報酬関数に渡されるので、ここでは数字だけの回答を指定し、正解の 77 との文字列一致を \(\{0, 1\}\) の二値報酬としています。
from trl import GRPOConfig, GRPOTrainer
trainer = GRPOTrainer(
model=sft_model,
processing_class=tokenizer,
reward_funcs=reward_func,
args=GRPOConfig(
output_dir="outputs/grpo",
num_generations=4,
per_device_train_batch_size=4,
loss_type="grpo",
importance_sampling_level="token",
scale_rewards="group",
beta=0.04,
epsilon=0.2,
use_bias_correction_kl=False,
),
train_dataset=train_dataset,
)
trainer.train()
num_generations は本書の \(G\)、beta と epsilon は \(\beta,\epsilon\) に対応します。scale_rewards="group" で式 (5.1) のグループ内正規化を行います。
5.3.2節 / 式 (5.10)〜(5.12)Dr. GRPO の実装
本書の式 (5.11) の目的関数は次の通りです。
GRPOTrainer で、loss_type="dr_grpo"、scale_rewards="none" を指定します。計算の要点は次の通りです。
import torch
grouped_rewards = rewards.view(-1, G)
advantages = grouped_rewards - grouped_rewards.mean(-1, keepdim=True)
advantages = advantages.reshape(-1, 1)
ratio = (per_token_logps - old_per_token_logps).exp()
clipped_ratio = ratio.clamp(1 - epsilon, 1 + epsilon)
per_token_loss = -torch.minimum(
ratio * advantages, clipped_ratio * advantages
)
loss = (per_token_loss * mask).sum() / (mask.size(0) * max_completion_length)
advantages は式 (5.10) の \(\widehat{A}^{\mathrm{Dr}}_i\)、ratio は式 (5.13) の \(w_{i,t}(\theta)\) に対応します。per_token_loss は式 (5.12) の \(\widehat{W}^{\mathrm{Dr}}_{i,t}(\theta)\) の符号を反転したものです。TRLでは式 (5.11) に対し、さらに固定値 max_completion_length で割ります。
from trl import GRPOConfig, GRPOTrainer
trainer = GRPOTrainer(
model=sft_model,
processing_class=tokenizer,
reward_funcs=reward_func,
args=GRPOConfig(
output_dir="outputs/dr-grpo",
num_generations=4,
per_device_train_batch_size=4,
loss_type="dr_grpo",
importance_sampling_level="token",
scale_rewards="none",
beta=0.0,
epsilon=0.2,
max_completion_length=512,
),
train_dataset=train_dataset,
)
trainer.train()
scale_rewards="none" で式 (5.10) のように標準偏差による除算をせず、beta=0.0 で KL 項を除外します。max_completion_length は最大生成長であり、損失の正規化にも使われます。
5.3.3節 / 式 (5.14)〜(5.20)DAPO の実装
本書の式 (5.16) の目的関数は次の通りです。TRLでは、その符号を反転した損失を最小化します。
GRPOTrainer で、loss_type="dapo" を指定します。1プロセス・勾配蓄積なしの場合の計算の要点は次の通りです。
import torch
grouped_rewards = rewards.view(-1, G)
advantages = (grouped_rewards - grouped_rewards.mean(-1, keepdim=True)) / (
grouped_rewards.std(-1, keepdim=True) + 1e-4
)
advantages = advantages.reshape(-1, 1)
ratio = (per_token_logps - old_per_token_logps).exp()
clipped_ratio = ratio.clamp(1 - epsilon, 1 + epsilon_high)
per_token_loss = -torch.minimum(
ratio * advantages, clipped_ratio * advantages
)
loss = (per_token_loss * mask).sum() / mask.sum().clamp(min=1)
advantages は式 (5.18) の \(\widehat{A}_{i,t}\)、clipped_ratio は式 (5.20) のクリッピングに対応します。per_token_loss は式 (5.19) の \(\widehat{W}^{\mathrm{DAPO}}_{i,t}(\theta)\) の符号を反転したもので、分母の mask.sum() は式 (5.16) の \(\sum_i|y_i|\) です。分散学習や勾配蓄積を使う場合は、バッチ全体の有効トークン数で正規化します。
from trl import GRPOConfig, GRPOTrainer
trainer = GRPOTrainer(
model=sft_model,
processing_class=tokenizer,
reward_funcs=reward_func,
args=GRPOConfig(
output_dir="outputs/dapo",
num_generations=4,
per_device_train_batch_size=4,
loss_type="dapo",
importance_sampling_level="token",
scale_rewards="group",
beta=0.0,
epsilon=0.2,
epsilon_high=0.28,
),
train_dataset=train_dataset,
)
trainer.train()
epsilon と epsilon_high は本書の \(\epsilon_{\mathrm{low}},\epsilon_{\mathrm{high}}\) に対応します。scale_rewards="group" で式 (5.18) のグループ内正規化を行い、beta=0.0 で KL 項を除外します。
この設定は損失の集計とクリップ幅に対応します。本書のDAPO全体を再現するには、式 (5.14)・(5.17) の動的サンプリングと、式 (5.15) の長さに応じた報酬を別途組み込む必要があります。
5.3.4節 / 式 (5.22)〜(5.25)GSPO の実装
GSPOも GRPOTrainer を使います。式 (5.22) の重要度比は次の通りです。
式 (5.24)・(5.25) の目的関数は次の通りです。TRLでは、その符号を反転した損失を最小化します。
importance_sampling_level="sequence" を指定すると、トークンごとの対数比を回答内で平均してから指数を取ります。GRPOの ratio に当たる計算を、次のように変更することに対応します。
log_ratio = per_token_logps - old_per_token_logps
sequence_log_ratio = (log_ratio * mask).sum(-1) / mask.sum(-1).clamp(min=1)
ratio = sequence_log_ratio.exp()
advantages = advantages.reshape(-1)
clipped_ratio = ratio.clamp(1 - epsilon, 1 + epsilon)
loss = -torch.minimum(ratio * advantages, clipped_ratio * advantages).mean()
ここでは ratio と advantages はそれぞれ回答ごとに1つの値を持ち、式 (5.22) の \(\xi_i(\theta)\) と式 (5.23) の \(\widehat{A}_i\) に対応します。最後の行が式 (5.24)・(5.25) の符号を反転した損失です。
from trl import GRPOConfig, GRPOTrainer
trainer = GRPOTrainer(
model=sft_model,
processing_class=tokenizer,
reward_funcs=reward_func,
args=GRPOConfig(
output_dir="outputs/gspo",
num_generations=4,
per_device_train_batch_size=4,
loss_type="grpo",
importance_sampling_level="sequence",
scale_rewards="group",
beta=0.0,
epsilon=0.0004,
),
train_dataset=train_dataset,
)
trainer.train()
式 (5.25) には KL 項がないため beta=0.0 とします。また、epsilon は式 (5.24) の \(\epsilon\) に対応します。