基盤モデルのための強化学習の実装

『基盤モデルのための強化学習』第4章・第5章の数式と、TRL の実装の対応を補足します。記号と各手法の導出は本書を参照してください。 なお、TRL の参照コミットは、2026年10月2日時点で最新版の 4b4c05aとします。ただし、PPO は実装が含まれる v1.8.0 を参照します。

4.4.2節 / 式 (4.8)RLHF:報酬モデルの学習

本書の式 (4.8) の損失関数は次の通りです。

\[ \mathcal{L}_{\mathrm{RM}}(\phi) =-\mathbb{E}_{(x,y_w,y_\ell)\sim\mathcal{D}_{\mathrm{HF}}} \left[\log\sigma\left(r_\phi(x,y_w)-r_\phi(x,y_\ell)\right)\right]. \tag{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) の目的関数を最大化します。

\[ \max_{\pi}\;\mathbb{E}_{x\sim\rho} \left[\mathbb{E}_{y\sim\pi(\cdot\mid x)}[r_\phi(x,y)] -\beta D_{\mathrm{KL}}[\pi(\cdot\mid x)\|\pi_{\mathrm{ref}}(\cdot\mid x)]\right]. \tag{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つのモデルを用意します。

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) の損失関数は次の通りです。

\[ \mathcal{L}_{\mathrm{DPO}}(\theta) =-\mathbb{E}_{(x,y_w,y_\ell)\sim\mathcal{D}_{\mathrm{HF}}} \left[\log\sigma\left( \beta\log\frac{\pi_\theta(y_w\mid x)}{\pi_{\mathrm{ref}}(y_w\mid x)} -\beta\log\frac{\pi_\theta(y_\ell\mid x)}{\pi_{\mathrm{ref}}(y_\ell\mid x)} \right)\right]. \tag{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では、その符号を反転した損失を最小化します。

\[ J_{\mathrm{GRPO}}(\theta)=\mathbb{E}_{\mathcal{D}_{\mathrm{VR}},G,\pi_{\theta_{\mathrm{old}}}} \left[\frac{1}{G}\sum_{i=1}^{G}\frac{1}{|y_i|}\sum_{t=1}^{|y_i|} \left(\widehat{W}_{i,t}(\theta)-\beta\widehat{D}_{\mathrm{KL}}^{i,t}[\pi_\theta\|\pi_{\mathrm{ref}}]\right)\right]. \tag{5.4} \]

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) の目的関数は次の通りです。

\[ J_{\mathrm{Dr.GRPO}}(\theta)=\mathbb{E}_{\mathcal{D}_{\mathrm{VR}},G,\pi_{\theta_{\mathrm{old}}}} \left[\frac{1}{G}\sum_{i=1}^{G}\sum_{t=1}^{|y_i|}\widehat{W}^{\mathrm{Dr}}_{i,t}(\theta)\right]. \tag{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では、その符号を反転した損失を最小化します。

\[ J_{\mathrm{DAPO}}(\theta)=\mathbb{E}_{\mathcal{D},G,\pi_{\theta_{\mathrm{old}}}} \left[\frac{1}{\sum_{i=1}^{G}|y_i|}\sum_{i=1}^{G}\sum_{t=1}^{|y_i|} \widehat{W}^{\mathrm{DAPO}}_{i,t}(\theta)\right]. \tag{5.16} \]

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) の重要度比は次の通りです。

\[ \xi_i(\theta)=\exp\left(\frac{1}{|y_i|}\sum_{t=1}^{|y_i|} \log\frac{\pi_\theta(y_{i,t}\mid x,y_{i,<t})}{\pi_{\theta_{\mathrm{old}}}(y_{i,t}\mid x,y_{i,<t})}\right). \tag{5.22} \]

式 (5.24)・(5.25) の目的関数は次の通りです。TRLでは、その符号を反転した損失を最小化します。

\[ \widehat{W}^{\mathrm{GSPO}}_i(\theta) :=\min\left\{\xi_i(\theta)\widehat{A}_i,\; \operatorname{clip}\left(\xi_i(\theta),1-\epsilon,1+\epsilon\right)\widehat{A}_i\right\}. \tag{5.24} \]
\[ J_{\mathrm{GSPO}}(\theta)=\mathbb{E}_{\mathcal{D}_{\mathrm{VR}},G,\pi_{\theta_{\mathrm{old}}}} \left[\frac{1}{G}\sum_{i=1}^{G}\widehat{W}^{\mathrm{GSPO}}_i(\theta)\right]. \tag{5.25} \]

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\) に対応します。