TT Lab
はじめる
学ぶ 学習パス コース

MiniMind — 小さな言語モデルを最初から最後まで自分で学習する

DPO で口調を変え、内容が守られたかも測る

TT Labで続きを見る

目標

MiniMindのtrain_dpo.pyの損失(答えトークンの対数確率の合計で計算する−logσ(β·Δ))を自分で実装し、SFTの基準モデルをポリシーと基準の2つとして読み込んで、100学習ステップDPOを行います。選好ペアは、事実は同じで口調だけが違います(chosenは…입니다.、rejectedは…이다.。韓国語の語尾で、順に「です。」と「だ。」にあたります)。口調が変わるかと、口調を無視した正答率が守られるかを、一緒に測ります。

なぜ重要なのか

DPOは、報酬モデルと強化学習なしで、選好ペアだけでモデルを押します。実装は10行ですが、損失が「差」だけを見るので、押し方を間違えると、損失はよく下がるのにモデルは壊れます。chosenとrejectedを一緒に引き下げて、差だけを開くという形で。このコースのモデルで、学習率を10倍に上げると、実際にそうなります。 そのため、選好学習は、損失1つで判断しません。選好の指標(取り分けておいたペアでchosenのほうがもっともらしい割合、望む口調で答える割合)と、元の能力の指標(事実の正答率)を、前後で一緒に測るのが、このラボの要点です。

ステップ

  1. ヘルパースクリプト(/root/mm/dpo/dpolib.py)に、選好ペアをバッチに変える関数と、答えトークンの対数確率の合計(seq_logp)を作り、基準モデルで最初のペアの2つの対数確率を書いてください(出力先: /root/mm/dpo/logp.json)。
  2. dpolib.pyにdpo_loss(정책 chosen, 정책 rejected, 기준 chosen, 기준 rejected, beta)(プレースホルダーは、順にポリシーのchosen・ポリシーのrejected・基準のchosen・基準のrejectedです)を実装してください。ステップ1のlogp.jsonに、ポリシー=基準のときの損失(loss_policy_equals_ref)も書いてください。
  3. DPOの前の基準モデルの指標(選好精度・丁寧語の割合・2つの平均対数確率)を書いてください(出力先: /root/mm/dpo/baseline.json)。
  4. β 0.1・lr 1e-5・ペア8個ずつで100学習ステップDPOを行い、ポリシーを保存してください(出力先: /root/mm/dpo/dpo.pth)。
  5. 同じ指標を、DPOのあとのポリシーで測って、書いてください(出力先: /root/mm/dpo/after.json)。
  6. SFTデータの先頭40個の質問で、口調を無視した正答率を、前後で測って書いてください(出力先: /root/mm/dpo/drift.json)。
  7. ## 손실 ## 무엇이 바뀌었나 ## 무엇을 못 하나の3つのセクションを書き、DPOのあとの丁寧語の割合とchosenの平均対数確率を入れてください(出力先: /root/mm/dpo/report.md)。韓国語の見出しは、順に「損失」「何が変わったのか」「何ができないのか」という意味です。

参考

文の対数確率は答えトークンの合計

ヘルパースクリプト(/root/mm/dpo/dpolib.py)に、batch(쌍들)(前半がchosen、後半がrejectedの入力とラベル。プレースホルダーはペアの集まりです)と、seq_logp(모델, 입력, 라벨)(文ごとの答えトークンのlog pの合計。プレースホルダーは、順にモデル・入力・ラベルです)を作り、基準モデル(/opt/mm/ref/sft.pth)で、dpo.jsonlの最初のペアの2つの値を、chosen・rejectedとして書いてください(出力先: /root/mm/dpo/logp.json)。

log_softmaxをかけたあと、torch.gatherで正解トークンの値だけを取り出し、ラベルが-100の位置は、マスクで0を掛けて除きます(gatherの前に-100を0に変えないと、インデックスエラーになります)。2つの答えは、語尾(「です」にあたる語尾と「だ」にあたる語尾)だけが違うので、その部分の確率の差が、そのまま2つの値の差になります。

DPO損失を式どおりに実装する

dpolib.pyにdpo_loss(policy_chosen, policy_rejected, ref_chosen, ref_rejected, beta)を作ってください。4つの引数は、文ごとの対数確率の合計(1次元テンソル)で、−logσ(β·[(π_c − ref_c) − (π_r − ref_r)])の平均を返します。基準モデルの最初の8ペアで、ポリシー=基準のときの値を、logp.jsonのloss_policy_equals_refに書いてください。

ポリシーが基準と同じなら、かっこの中が0なので、−log σ(0) = ln 2です。採点ツールは、作成したdpo_lossを呼び出して、ランダムな6ペア・βが2つの場合に、論文の式と同じか、そしてポリシーがchosenをより好むとき、損失がln 2より小さいか(符号)を見ます。F.logsigmoidを使うと、数値が安定します。

DPOの前の指標を測る

スクリプト(/root/mm/dpo/metrics.py、引数: 重みのパス、結果のパス)を書いて、基準モデルで、dpo_val.jsonlの100ペアの選好精度(pref_acc)、先頭40ペアの質問の丁寧語の割合(polite_rate)、chosenとrejectedの平均対数確率(mean_logp_chosen・mean_logp_rejected)を測って、書いてください(出力先: /root/mm/dpo/baseline.json)。

SFTデータは、口調が半々だったので、基準モデルは、2つの口調を同じくらい混ぜて使います。指標を測るスクリプトを、重みのパスと結果のパスを引数で受け取るようにしておけば、ステップ5でそのまま使えます。

100学習ステップのDPO

スクリプト(/root/mm/dpo/train_dpo.py)を書いて、SFTモデルをポリシーと基準の2つとして読み込み、基準を凍結して、dpo.jsonlからペアを8個ずつ選び、β 0.1・lr 1e-5(AdamW)・100学習ステップ(シード42)でDPOを行ったあと、ポリシーのstate_dictを保存してください(出力先: /root/mm/dpo/dpo.pth)。

基準モデルの対数確率は、torch.no_grad()の中で求めます。MiniMindのデフォルトの学習率は4e-8ですが、この小さなモデルと短い学習では、1e-5が口調を移しながら、内容を守ります。損失はln 2から始まって下がります。

DPOのあとの指標を測る

ステップ3と同じ指標をdpo.pthで測って、書いてください(出力先: /root/mm/dpo/after.json)。採点ツールは、選好精度が0.8以上で、丁寧語の割合がDPOの前より20ポイント以上上がったかを見ます。

平均対数確率2つも、前後で比べてみてください。chosenとrejectedがどちらも下がるのは、DPOでよくあることです。問題は、chosenが下がりすぎて、モデルがどちらでもない答えを出すときです。

内容は守られたかを確かめる

sft.jsonlの1往復の対話の先頭40個の質問に、グリーディ生成で答えさせ、終わりの입니다.・이다.を取り除いて(韓国語で、順に「です。」と「だ。」にあたる語尾です)、データの答えと同じになる割合を、基準モデル(fact_acc_ref)とDPOのポリシー(fact_acc_dpo)で測って、書いてください(出力先: /root/mm/dpo/drift.json)。

口調を取り除いて比べれば、「何を言ったか」だけが残ります。選好ペアは、事実が同じで口調だけが違うので、うまくいったDPOなら、この値はほとんどそのままのはずです。10ポイント以上下がったら、基準から離れすぎています。

選好学習の効果と限界を残す

## 손실 ## 무엇이 바뀌었나 ## 무엇을 못 하나の3つのセクションを書き、ステップ5のpolite_rateとmean_logp_chosenを数字で入れてください(出力先: /root/mm/dpo/report.md)。韓国語の見出しは、順に「損失」「何が変わったのか」「何ができないのか」という意味です。

最後のセクションには、選好ペアで教えられないもの(モデルが知らない事実)と、学習率を上げたときに見たことを書けばよいです。