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

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

モデルの出来を数字ひとつにする

TT Labで続きを見る

一言でいうと

モデルが正解トークンに与えた確率の対数の符号を反転して、平均したものが交差エントロピーで、それを指数で戻したものがパープレキシティです。パープレキシティは、「各位置で平均いくつの候補に迷っているか」と読みます。

なぜ必要なのか

学習も評価も、結局は1つの数字を見て動きます。損失が下がれば回し続け、下がらなければ何かを変えます。ところが、その数字を作る過程には、エラーを出さずに静かに間違う箇所が何か所もあります。

実際に出会う姿は、次のようなものです。損失がある時点でinfになり、次のバッチからnanが広がります。あるいは、損失曲線はきちんと下がっているのに、生成された文章はひどい出来です。あるいは、昨日測ったパープレキシティと今日測った値が違うのに、モデルはそのままです。あるいは、他人の論文の数字と自分の数字が2倍違うのに、どちらが間違っているのかわかりません。

原因は、たいてい次の4つのうちの1つです。対数をいつとったか、ラベルを1つずらしたか、パディングの位置を除いたか、対数の底が何かです。4つとも、コードでは1行で、間違っても例外は出ません。そのため、自分で測ってみるまでは見えません。

対数をいつとるのか

ソフトマックスは、スコアの行を、合計が1の分布に変えます。私たちに必要なのは、その確率ではなく、対数確率です。それなら、「確率を求めたあとで対数をとればよいのでは」と思うのですが、その順序が、下限で崩れます。

倍精度の実数が持てる最も小さい正の数は、5e-324付近です。正解トークンのスコアが、残りよりかなり下にあれば、その確率はそれより下に落ちて、0.0に沈みます。0の対数はありません。Pythonのmath.logは、その場で例外を投げ、例外を避けるために-infを入れておくと、その位置の損失が+infになって、平均をとった瞬間に、文全体がinfに染まります。

答えは、確率をそもそも作らないことです。

# log p_i = x_i - (max + log sum exp(x - max))
top = max(xs)
lse = top + math.log(sum(math.exp(x - top) for x in xs))
logp = [x - lse for x in xs]

ここには、割り算がありません。大きい値を引いてからexpするのであふれず、対数確率を引き算で得るので、下限にも達しません。確率がどんなに小さくても、その対数は、単に小さい負の数にすぎません。フレームワークが、ソフトマックスと損失を1つの関数にまとめて売っている理由が、これです。2つの演算を別々に呼ぶと、その間で情報が失われます。

既存のラボで扱ったソフトマックスの安定化は、上側を防ぐものでした(expがあふれること)。ここで防ぐのは、下側です。同じmaxを引くことが2つの仕事をしていますが、崩れる場所も症状も違います。

1つの位置から文へ

1つの位置の損失は、1行です。

損失(t) = -log p(正解トークンt)

正解に確率1を与えていれば0で、確率が小さくなるほど大きくなります。他のトークンに何を与えたかは、別に数えません。合計が1なので、正解の取り分が、そのまま残りの取り分だからです。

文の損失は、それらの値の平均です。合計ではなく平均なのは、長さの異なる文を比べるためです。合計で測ると、長い文は常に悪い文になります。statisticsの平均が行うことと同じですが、何を分母に入れるかが、あとで問題になります。

パープレキシティは、目盛りを変えただけ

損失1.06という値は、ピンときません。指数で戻すと、読める数字になります。

パープレキシティ = exp(平均損失)

目盛りをつかむ方法が1つあります。何も知らないモデルを入れてみることです。語彙全体に同じスコアを与えると、確率は1/Vで、損失はlog Vなので、パープレキシティは、ちょうどV、つまり語彙サイズになります。そのため、パープレキシティが語彙サイズの近くなら、そのモデルは何も学べていないということで、それより大きければ、一様分布にも劣るということです。

ここからすぐに出てくる結論が、1つあります。語彙が異なる2つのモデルのパープレキシティは、比べられません。目盛りの出発点が違うからです。トークナイザーが違えば、同じ文章を分ける断片の数も違うので、分母まで違います。論文の数字を自分の数字と並べる前に、語彙とトークナイザーが同じかをまず見る必要がある理由です。

ラベルは1つずれている

言語モデルは、位置tの出力で、位置t+1のトークンを当てます。Attention Is All You Needのデコーダーが行うことがそれで、Hugging Faceの生成のドキュメントが説明する次のトークンの予測も、同じ話です。

そのため、損失を測るときに、ロジットは後ろを1つ捨て、トークンは前を1つ捨てます。最後の行には、当てるべき次のトークンがなく、最初のトークンには、それを予測した行がないからです。

1つずらさないと、どうなるでしょうか。エラーは出ません。モデルに「いま見ているトークンを当てなさい」と命じたことになるので、損失が悪く出るだけです。逆に、ずらすべき場所を2回ずらしたり、向きを逆にずらしたりすると、数字が不自然によくなることもあります。自分がすでに見たものを答えとして渡すことになるからです。損失曲線はもっともらしいのに、生成結果がひどいなら、まずここを見ます。

パディングはタダのスコア

バッチにまとめるために、短い行を埋めて長さをそろえます。その埋めのトークンは、内容ではなく、位置の目印です。

問題は、埋めのトークンが、当てるのがあまりにも簡単だという点にあります。文が終わったあとは、いつも同じものが来るので、モデルがすぐに確信します。その位置を平均に入れると、損失が下がり、埋めが多いバッチほど、さらに下がります。モデルはそのままなのに、バッチの構成を変えるだけで、数字がよくなります。

除く場所は、2か所です。足す側でも除き、割る側でも除きます。割る側を忘れると、残った位置の損失を全体の長さで割ることになって、値が静かに小さくなります。このミスは、特に見つけにくいです。方向がいつも「よくなる」ほうなので、疑うきっかけがないからです。

ビットで測ると

対数の底を2に変えると、単位がビット/トークンになります。割り算1回です。

ビット/トークン = 自然対数の損失 / ln 2

底を変えただけなので、2 ** 비트はexp(자연로그 손실)と同じ値です(プレースホルダーはビット/トークンの値と、自然対数の損失です)。同じものを別の物差しで読んだだけなのに、圧縮の分野の文献はビットで書き、ディープラーニングの分野は自然対数で書くことが多いので、数字だけを見ると、2倍近く違うように見えます。他人の表を書き写す前に、底を確認する必要があります。

現場での姿

第1に、損失が突然infやnanになります。確率を先に作ってから対数をとったコードで、非常に低い確率に出会った瞬間です。対数確率を引き算で得るように直せば、なくなります。

第2に、損失は下がるのに、生成物が悪いです。ラベルのシフトがずれていないかを、まず見ます。答えを先に見せて当てさせれば、損失はいくらでも下がります。

第3に、同じモデルのパープレキシティが、実行ごとに違います。評価バッチのパディングの割合が変わった可能性が高いです。マスキングを正しく行えば、バッチの構成が変わっても、値は揺らぎません。

第4に、他人の数字と2倍違います。対数の底、語彙サイズ、トークナイザー、分母をトークンにしたのか単語にしたのかを、順に合わせてみます。たいてい、そのうちの1つです。

第5に、損失1つだけを見てデプロイして、事故が起きます。パープレキシティは、「次のトークンをどれだけうまく当てるか」であって、「役に立つ答えをするか」ではありません。下がっていることを確認する用途にはよいのですが、それだけで、よいモデルだと言うことはできません。

実務で本当に大切なこと

次のラボですること

/root/work/tf-loss/loss.pyを、1ステップずつ育てていきます。実際のモデルを呼び出すのではなく、同じ計算を標準ライブラリだけで自分で作ります。このPodのシステムのPythonには、numpyもtorchもtransformersもなく、numpyは/opt/onnx-lab/bin/pythonの中にしかありません。そのため、ここに出てくる数字は、すべて自分が作ったデータで測ったものです。

安定したlog-softmaxから始めます。そのあと、わざと間違った順序の版を別に作って、確率を先に求めたときに、-infが実際に出ることを目で見ます。2つの関数が並んでいて初めて、どの箇所で分かれるのかが見えます。

そこから、1つの位置の損失、複数の位置の平均、指数で戻したパープレキシティへと上がっていきます。何も知らないモデルのパープレキシティが、語彙サイズと同じになることも、自分で確認します。

最後の3つのステップが、このラボの要点です。同じデータを使って、ラベルをずらしたときとずらさないとき、パディングを除いたときと除かないときの数字を、並べます。4つの数字がすべてエラーなく出て、4つともそれらしく見えるのに、互いに違うということ。それが、このモジュールが見せたいことのすべてです。最後に、底が2の対数に移して、ビット/トークンでも読みます。採点ツールは、自分で作ったモジュールを実際に呼び出し、毎回異なるロジットで関数を直接叩いて、採点ツールが別に計算した値と照合します。