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

Transformer — アテンションを手で計算する

損失をひとつ手で作る

TT Labで続きを見る

目標

ロジットから出発して、「このモデルがどれだけうまくやっているか」を1つの数字にする過程を、標準ライブラリだけで自分で組みます。安定したlog-softmax、1つの位置の損失、複数の位置の平均である交差エントロピー、指数で戻したパープレキシティまで進めたあと、ラベルを1つずらすことと、パディングの位置を除くことが、その数字をどれだけ変えるかを、並べて測ります。最後に、底が2の対数に移して、ビット/トークンでも読みます。

なぜ重要なのか

学習も評価も、この数字1つを見て動きます。ところが、この数字を作る過程には、エラーを出さずに静かに間違う箇所が4つあります。対数をいつとったか、ラベルを1つずらしたか、パディングを除いたか、対数の底が何か。4つとも、コードでは1行で、間違っても例外は出ず、たいていは数字がよくなる方向に間違います。そのため、疑うきっかけがありません。 このラボは、実際のモデルを呼び出しません。このPodのシステムのPythonには、numpy・torch・transformersがありません。その代わりに、位置ごとに語彙全体に対するスコアの行を決定的に作っておき、その上で同じ計算を手で組みます。そのため、「あるモデルのパープレキシティはいくつ」といった話は、ここではしません。出てくる数字は、すべて自分が作ったデータで測ったものです。 隣のモジュールが、分布から1つを選ぶ方法(温度・top-k・top-p)なら、ここは、その分布がどれだけ間違っているかを測る方法です。選ぶ前に測ることが先です。 採点ツールは、書いておいた説明を信用しません。自分で作ったモジュールを実際に呼び出し、毎回異なるロジットで関数を直接叩いて、採点ツールが別に計算した値と照合します。入力は実行のたびに変わるので、値を暗記して入れることはできません。

ステップ

  1. /root/work/tf-loss/loss.pyにVOCAB・PAD_ID・SEQ・dataset()と、log_softmax(xs)を作成してください。確率を経由せず、ロジットから直接対数確率に進みます。
  2. NEG_INFとnaive_log_softmax(xs)を追加して、わざと間違った順序の版を作成してください。確率を先に求めてから対数をとり、下限で-infが出ることを再現します。
  3. token_loss(logits, target)を追加して、1つの位置の損失を測るようにしてください。正解トークンの対数確率の符号を反転した値です。
  4. cross_entropy(rows, targets)を追加して、複数の位置の損失を平均するようにしてください。
  5. perplexity(rows, targets)とuniform_perplexity(vocab_size, length)を作成してください。一様分布で、パープレキシティが語彙サイズと同じになることを確認します。
  6. shift_pairs(rows, ids)・shifted_loss(rows, ids)・unshifted_loss(rows, ids)を作成して、ラベルを1つずらしたときとずらさないときを、並べて測ってください。
  7. kept_positions(targets, pad_id)とmasked_cross_entropy(rows, targets, pad_id)を作成して、パディングの位置を除いて測るようにしてください。分子と分母の両方で除きます。
  8. bits_per_token(loss)・nats_per_token(bits)を追加し、/root/work/tf-loss/loss_report.jsonと、/root/work/tf-loss/loss_report.mdに、測った値を記録してください。

参考

ロジットから直接対数確率へ進む

/root/work/tf-loss/loss.pyに、VOCAB(12以上)・PAD_ID・SEQ(16個以上、後ろの3個以上がPAD_ID)・dataset()と、log_softmax(xs)を作成してください。dataset()は、(로짓 줄 목록, 토큰 번호 목록)(プレースホルダーはロジットの行のリストとトークン番号のリストです)を返し、乱数を使いません。log_softmaxは、確率を経由せず、ロジットから直接対数確率を出します。

mkdir -p /root/work/tf-loss。式は、x_i - (max + log sum exp(x - max))の1行です。割り算がない点が要点です。確率を作って割ってから対数をとると、非常に小さい確率が0.0になって、対数が崩れます。log_softmax([0.0, -800.0])の2番目の要素が有限の値(-800の近く)なら、正しくできています。dataset()は、位置tの行が、位置t+1のトークンに最も高いスコアを与えるように作り、パディングが正解の位置には、スコアをさらに大きく載せてください。

わざと崩して確かめる

NEG_INFとnaive_log_softmax(xs)を追加してください。今度は確率を先に求めてから対数をとります。確率が0.0に沈んだ位置は、math.logが例外を投げるので、自分でNEG_INFで埋めます。同じ入力で、log_softmaxは有限なのに、こちらだけ-infが出ることを確認してください。

NEG_INF = float("-inf")です。順序を変えるだけで済みます。expした値を合計で割って確率を作り、その確率の対数をとります。中間の値では、前のステップの関数と同じ答えが出て、下限でだけ分かれます。[0.0, -800.0]のように、差の大きい行を入れてみてください。確率が0.0かどうかを見て避ける必要があり、math.logをそのまま呼ぶと、例外で終わります。

1つの位置の損失を測る

token_loss(logits, target)を追加してください。正解トークンにモデルが与えた対数確率の符号を反転した値です。正解に確率1を与えていれば0で、確率が小さくなるほど大きくなります。

1行です。-log_softmax(logits)[target]。符号を反転するのを忘れると、値がすべて負になって、「損失が下がる」という言葉が逆になります。ロジットそのものを使ってはいけません。他のトークンに何を与えたかは、別に数えません。合計が1なので、正解の取り分が、そのまま残りの取り分です。

複数の位置の平均をとる

cross_entropy(rows, targets)を追加してください。位置ごとにtoken_lossを出して、平均を返します。targetsが空なら0.0です。

合計ではなく平均です。合計で測ると、長い文がいつも悪い文になって、長さの異なる文章を比べられません。zip(rows, targets)でペアにして足してから、len(targets)で割ってください。ここで何を分母に入れるかが、ステップ7で再び問題になります。

パープレキシティの目盛りをつかむ

perplexity(rows, targets)とuniform_perplexity(vocab_size, length)を作成してください。前者はexp(평균 손실)(プレースホルダーは平均損失です)で、後者は、すべてのスコアが同じロジットの行を作って、パープレキシティを測ります。結果がvocab_sizeと同じになるかを確認してください。

math.exp(cross_entropy(rows, targets))の1行です。位置ごとにexpをとって平均することと混同しやすいのですが、一様分布では2つの値が偶然同じになるので、そのテストでは区別できません。uniform_perplexityは、[[0.0] * vocab_size] * lengthの形の行を作って、適当な正解を入れれば済みます。すべての確率が1/Vなので、損失はlog V、指数をとるとVです。この値が、パープレキシティの目盛りの出発点です。

ラベルを1つずらす

shift_pairs(rows, ids)・shifted_loss(rows, ids)・unshifted_loss(rows, ids)を作成してください。shift_pairsは、(rows[:-1], ids[1:])で、位置tのロジットが位置t+1のトークンを当てるようにペアを合わせます。残りの2つは、ずらした版とずらさない版の交差エントロピーです。

最後の行には当てるべき次のトークンがなく、最初のトークンにはそれを予測した行がありません。そのため、ロジットは後ろを、トークンは前を、1つずつ捨てます。向きを逆にずらすと(rows[1:]、ids[:-1])、すでに見たものを答えとして渡す形になって、数字が不自然によくなります。自分のデータでは、ずらしたほうの損失が、ずらさないほうより確実に小さい必要があります。dataset()が、位置t+1のトークンにスコアを載せておいたからです。

パディングを除いて測る

kept_positions(targets, pad_id)とmasked_cross_entropy(rows, targets, pad_id)を作成してください。前者は、正解がパディングではない位置の番号のリストで、後者は、その位置だけを足してその位置の数で割った平均です。残る位置がなければ0.0です。

除く場所が2か所あります。足す側と、割る側です。割る側を忘れて全体の長さで割ると、値が静かに小さくなり、方向がいつもよくなるほうなので、疑うきっかけがありません。kept_positionsを別に置くと、何が残ったかを目で見られます。ずらしたペアにこの関数を使うと、パディングを入れて測った値より大きくなる必要があります。パディングは、当てやすい位置だからです。

4つの数字を並べる

bits_per_token(loss)とnats_per_token(bits)を追加し、/root/work/tf-loss/loss_report.jsonにはvocab_size・pad_id・seq_len・pad_count・kept・dropped・unshifted_loss・shifted_loss・masked_loss・unshifted_perplexity・shifted_perplexity・masked_perplexity・bits_per_token・uniform_perplexity・probe_gap・naive_is_inf・stable_logprobを、/root/work/tf-loss/loss_report.mdには## 무엇을 쟀나 ## 한 칸 어긋나면 ## 패딩을 빼면 ## 비트로 재면の4つの節で書いてください(韓国語の見出しは順に、「何を測ったか」「1つずれると」「パディングを除くと」「ビットで測ると」という意味です)。

数字は手で書かず、自分のコードを実際に動かして得た値で埋めてください。masked_lossは、shift_pairsでずらしたペアにmasked_cross_entropyを使った値で、shifted_lossは、同じペアをマスキングなしで測った値です。bits_per_tokenは、masked_lossを基準に出します。2 ** 그 값(プレースホルダーはその値です)がmasked_perplexityと同じかを、自分で確認してみてください。uniform_perplexityは、uniform_perplexity(VOCAB, 8)です。probe_gapは800で固定で、naive_is_infは、naive_log_softmax([0.0, -800.0])[1]が-infかどうか、stable_logprobは、同じ入力でlog_softmaxが出した2番目の要素です。