確率 入門
Transformer が実際に走っている確率を、たった 1 枚の絵の上に組み上げる:高さの合計が 1 になる棒の並び。分布、期待値と分散、ベイズ、尤度、サンプリング、エントロピーと交差エントロピー —— このページのすべての数値は隣の図が計算しているので、どこまで動かしても主張は崩れない。
確率とは取り分のこと
0 と 1 のあいだの1 つの数。このページの残りは、すべてその数についての算術でしかない。
確率と呼ばれるものは 2 種類ある。1 つは頻度:実験を 10 万回まわして数える。もう 1 つは 1 度きりの出来事に対する信念の度合い。同じ 3 つの規則に従うから同じ名前で呼ばれるのであって、絵に描けるのは頻度のほうだ。だから描く —— 100 個のマス、1 回の試行に 1 個、塗られたマスが実際に起きた回数:
最後のマスの先には、スライダーの行き場がないことに注目。確率が [0, 1] にいるのは、それが何かの取り分であり、取り分が全体を超えられないからだ。では右側の読み値が分数をやめて必ずと言い、では決してと言う。
頻度という読み方には警告が付いてくる。公平なコインを投げると、表の割合はしばらくさまよってから確率に落ち着く —— しかも多くの人が思うよりずっと遅い。投げる回数を右へ押し出して、揺れが狭まる様子を見てほしい:
1 回投げた時点で割合は 1.000。コインは表を出し、ここまでの証拠は「このコインは必ず表」と言っている。割合を 0.570 まで引き寄せるのに、 0.501 にするのに 10,000 回かかる。帯は標準誤差 1 つぶん0.5/√n —— データ 4 倍で誤差が半分だ。
3 つ目の規則は、以降のすべての節が寄りかかる。各結果に 1 本の帯の取り分を与える ——雨、曇、晴がちょうど敷き詰める —— ので、境界を動かしても隣から幅を奪うだけだ:
読み値の右端を見てほしい。晴は選ばれていない。計算されている、1 − 0.35 − 0.36 と。これが合計が 1 の実務上の意味だ —— どれか 1 つの数は自由ではない。 0.35、0.36、0.40 という予報は楽観ではなく、帳尻の合わない算術だ。
最初から合計が 1 のものはほとんどないので合計で割る —— うまくいかなくなるまでは。3 つ目のスコアを 0 より下へドラッグしてほしい:
c = 3 ではスコア 3, 4, 3 が合計 10 のもとで 0.30, 0.40, 0.30 に正規化される。 では合計が 6 になり、3 つ目の取り分は −0.17。数は 3 つ返るし、例外も出ないし、形は無意味だ。正規化は何も検査しない —— スコアが全部 0 なら、何も言わずに NaN を返す。
分布とは形のこと
結果 1 つにつき数 1 つ。どれも負ではなく、合計はちょうど 1。
§01 の帯を結果の数だけ切り分け、その一片ずつを高さとして立てる。機械学習のあらゆるプロットは、その絵の変奏にすぎない。ここの破線は公平なサイコロが各目に与える高さだ —— 重りを増やして、6 の目が残り 5 面から余分な確率質量を奪う様子を見てほしい:
注目してほしい。読み値の合計はびくともしない。にすると6 の目が 1.00、残り 5 面はちょうど 0 —— それでもまだ合法な分布で、ただ偶然が残っていないだけだ。この退化した隅は後で効いてくる。温度ゼロも argmax も one-hot ラベルも、すべてこの同じ形をしている。
モデルが出すのは確率ではない。出すのは logits —— トークンごとに 1 つ、上下限のない実数スコア —— で、softmax がそれを指数にとって合計で割る。2 番目のトークンの logitをドラッグして、下の行が追従するのを見てほしい:
exp は単調なので上下の行の順序は同じであり、 softmax が argmax を変えることはない。変わるのは間隔だ。手にしているトークンが上下同時に動く。 logit の差 1 は確率の比 e ≈ 2.72 になるので、先頭のトークンと末尾のトークンの差 5 は 148 倍 —— 0.577 に対して 0.0039 になる。
実装は必ず、指数をとる前に最大の logit を引く。分子と分母で打ち消し合うので答えは同一で、得られるのは「float64 では exp(800) がInfinity で exp(0) は 1」という事実だけだ。 50,257 トークンの語彙でも、この一巡はトークンあたり指数 1 回と除算 1 回。 logits を作った行列積に比べれば無に等しい。一方、logits を先に温度で割るのは本当の変更になる —— 温度をゼロへ引いて、先頭のトークンが行を食い尽くす様子を見てほしい:
T が まで下がった瞬間、先頭のトークンは小数 3 桁で 1.000 —— 貪欲デコードであり、 API の temperature=0 の意味そのものだ。 では 0.278 まで下がり、1/5 へ向かう。温度が変えるのは信念ではなく、そのどれだけがサンプルに残るかだ。
連続の結果にはもう 1 つ考えが要る。そしてそれが人をつまずかせる。密度は確率ではない。x の単位あたりの確率なので、 1 より大きくてかまわない。ベルを細くして、窓を沿って動かしてほしい:
見てのとおり σ = 1 では頂点が 0.399、窓の中身は 0.383。σ = 0.2 では頂点が 1.995 —— 確率が決して取れない値 —— なのに同じ窓の中身は 0.988 で、これは何の問題もない。面積が確率であり、高さはその速さでしかない。密度が 1 を超えるのは不具合報告ではない。
どこにいて、どれだけ広がるか
分布ぜんぶを代弁する 2 つの数 ——釣り合いの点と距離の 2 乗。
結果を直線上に並べ、各結果の確率をその上に重りとして吊るすと、分布全体に釣り合いの点ができる。確率質量を1 つの目へ寄せて、支点が梁の上を追ってくるのを見てほしい:
なぜならどの目も等しく起こりうるからで、公平なサイコロは 3.50 で釣り合う —— 決して出ない値だ。期待値について最初に受け入れるべきはこれで、それは重み付きの住所であって、結果ではない。まで入れると支点は 6.00 に届き、同じ重りを 1 に入れれば 1.00 まで引きずられる。
広がりには 2 つ目の数が要る。そしてそれを機能させるのが 2 乗だ。どの結果も平均からの距離の 2 乗を、自分の確率で重み付けして支払う。確率質量を両端へ押しやって、請求額が育つのを見てほしい:
真ん中の目は最も起こりやすいときでさえ、ほとんど寄与しない。半目ぶん離れたところの2 乗項は 0.25、両端は 6.25 だ。この 2 乗は飾りではない。独立な変数のあいだで分散が足し合わせられるのはこれのおかげで、絶対距離ではそうならない。
2 乗の代償は単位だ。分散 4.15の単位は目の 2 乗で、誰にも思い描けない。平方根をとれば、この数は本来いるべき軸の上へ戻ってくる:
注目してほしい。一様のときの σ は 1.71 で、区間 μ ± σ は質量の 0.67 を覆う —— 6 面のうち 4 面だ。広がりをまで押すと σ は 2.38 に育つのに、覆われる質量は 0.13 まで落ちる。 σ は物差しであって、容器ではない。
もう 1 つ。これは実際に人がやる間違いだ。分散を推定するとき2 乗偏差を n で割ると、答えは小さく出る。標本数をドラッグしてほしい:
見てのとおり n = 2 では ÷n の推定量が平均 1.458 を返す。真の値 2.917 のちょうど半分で、しかも何も言わない。 で真値の 90.0%、n = 20 で 95.0%。偏りは (n−1)/n で、小さくはなるが閉じない。n − 1 で割るほうはどの n でも線に乗る。
標本平均の分散は σ²/n なので、標準誤差は §01 の 1/√n —— バッチサイズの理屈もこれだ。 GPT-3 の訓練バッチは 320 万トークン。 1000 倍にして、勾配は 32 倍しか静かにならない。
別のことが起きたという前提のもとで
条件付けとは正規化のやり直しだ —— 証拠が排除した人を消して、残りを拡大し直す。
2 つの事象は 1 つの正方形に収まる。左右を罹患の有無で、上下を検査結果で切ると、各タイルの面積がそのまま同時確率になる。ここでの検査は罹患者の 99% を捕まえる。縦の切れ目をドラッグして、その病気の多さを変えてほしい:
有病率 10% では「罹患かつ陽性」のタイルが 0.099、「健康なのに陽性」のタイルが 0.090 —— 探している相手には 99% 正しい検査なのに、面積はほぼ同じだ。健康の列は 9 倍広いので、誤り率が小さくても板は同じくらい厚くなる。
陽性という結果は陰性の行をまるごと消し、タイル 1 つと板 1 枚だけが残る。残ったものはまだ分布ではない —— 合計は 1 ではなく 0.189 だ —— ので、分布になるまで引き伸ばす。下の帯の右端をステージいっぱいまでドラッグしてほしい:
ドラッグしているあいだ、帯の中の比率が一度も変わらないことに注目。P(罹患 | +) = P(罹患 ∧ +) / P(+) は両方を同じ数で割るので比率は固定され、動くのは「全体」という語の意味だけだ。答えはドラッグの両端で 0.524 と読める。
次に病気をまれにする。検診プログラムが実際にやっていることだ。検査を両方向とも 99% に保ったまま、開始時の 10 分の 1 から有病率を対数軸の下へ引いてほしい:
有病率 では答えは 0.090。99% 正確な検査で、陽性で、それでもあなたが健康である確率は 91% —— 真陽性 1 つに対して偽陽性が 10.1 個来るからだ。釣り合いは有病率ちょうど にあり、そこで 2 つの誤り率が相殺する。それより下では基準率が勝ち、検査精度を上げても交点が動くだけだ。
これはこのページで唯一、静かだからこそ危険な失敗だ。例外も出ず、どの数字も間違って見えず、同じ検査の P(+ | 罹患) と P(罹患 | +) が 0.99 と 0.09 になる。独立とは、この算術が要らない特別な場合のことで —— それにも絵がある。2 本の切れ目をずらしてほしい:
ずれが 0 のとき P(A | B) = P(A) = 0.50 で、正方形はただの格子だ。 まで押すとP(A | B) は 0.80 になり、P(A) は動かない —— 周辺確率は作りつけで固定されているので、 2 本の切れ目のあいだの段差が従属性のすべてになる。独立とはその段差が 0 だという主張であって、それ以上ゆるい意味はない。
どのモデルがこれを起こしやすくしたか
問いを裏返す。パラメータが何を予測するかではなく、見えたものを最も起こりやすくするのはどのパラメータか。
コインを 20 回投げて表が 13 回出たとしよう。ありうる偏りはどれもこの結果に確率を与えるので、データを固定して偏りの関数として読んだその曲線が尤度だ。偏りを軸に沿ってドラッグしてほしい:
頂点に注目 —— 0.65、ちょうど 13/20 だ。コインでは最大尤度推定量が標本頻度そのものだからだ。 θ = 0.50 では曲線は頂点の 0.40 —— 「公平」は競り負けただけで、否定されてはいない。データをまで押すと頂点は θ = 1.00 へ跳ぶ。裏は出ないと言い切るモデルだ。
この曲線には不都合が 2 つある。偏りについての分布ではないこと、そして 20 回の時点で値がもう微小なことだ。対数をとってほしい:
log は単調なので頂点は動かない。 L を最大にするものが log L も最大にする。しかも値が読める —— 頂点で −12.95、公平で −13.86 —— 生の尤度は 2.4e−6 と 9.5e−7 と出る。
これは趣味の問題ではない。独立な確率を掛け合わせていくと、系列が面白くなるはるか手前で積が float64 を使い果たす。言語モデルが満足する程度のトークンあたり確率をとって、項数を 1 から外へ押してほしい:
見てのとおり、トークンあたり p = 0.02 だと、で積は 1.3e−170、でちょうど 0 になる —— float64 の最小の非正規化数は 4.94e−324 だ。対数の和は −747.2 と読め、そのまま下り続ける。例外は出ない。積が 0 になり、log 0 が −∞ になり、最初の NaN が下流の勾配で顔を出す。
この対数の符号を裏返せば、言語モデルが訓練に使う損失になる。曲線は 1 本、読む場所も 1 か所 ——次に来たトークンにモデルが与えた確率だ:
p = 1 では損失は 0。唯一ただで済む答えだ。 0.5 では 0.693 nat、0.2 では 1.609、0.002 では 6.215。損失に上限はないので、自信をもって間違えることは無制限に罰せられ、自信のなさは安く済む —— この非対称が訓練信号のすべてだ。実際に現れたトークンにちょうど 0 を割り当てたモデルが無限大の損失と死んだ実行を生むのも、同じ理由による。
そこから 1 つ引く
分布は記述する。サンプラーは実行する —— しかも必要な乱数はたった 1 つだ。
5 つの確率を、長さ 1 の線の上に端から端まで並べる。生成器に[0, 1) の一様乱数 u を 1 つもらい、帯の上に落とす。落ちたブロックがそのままトークンだ。ダーツをドラッグしてほしい:
注目してほしい。これが逆変換サンプリングで、アルゴリズムはこれで全部だ。ブロックの境界は累積和 ——ダーツのいるブロックは帯の横に0.000 → 0.577と書かれている —— なので、そこを二分探索すれば 1 つの一様乱数が O(log V) で 1 つのトークンになる。どのブロックも自分の幅どおりの頻度で当たる。必要な性質はそれだけだ。
サンプリングは、総和を取るには大きすぎるものを測る手段でもある。 n 個引いて、数えて、割る。答えの良さは n だけで決まる。サンプル数を外へ押してほしい:
両軸とも対数なので、この直線そのものが主張になっている。サンプル数が 2 桁上がって、区間がやっと 1 桁下がる。 10 サンプルで ±0.31、 で ±0.031、 ±0.01 には 9,604 個。モンテカルロは次元数を気にしない —— 気にするのは 1/√n だ。
放っておけばサンプラーはいずれ裾に手を伸ばすし、言語モデルの裾はたいてい無意味なので、デコーダはそこを切る。残す個数を下げて、捨てられる確率質量が現れるのを見てほしい:
見てのとおり では残る質量が 0.913 で、語彙の半分が消えている。 では 0.448 になり、引くものが残らない —— 反対側から来た貪欲デコードだ。核サンプリングは操作を裏返す。残したい質量を 0.9 と先に決め、 k はそれに要るだけ —— ここでは 5 —— にさせる。
裾がそもそも問題になるかどうかは、人があまり計算しない数にかかっている。 1000 分の 1 のトークンはまれだが、生成は長い。長さを対数軸の外へ押してほしい:
1 トークンなら確率は宣伝どおり 0.001。 なら 0.394、 なら 0.632。読み値の積 qn が 1 に届いた瞬間に1 − (1−q)ⁿ は 1 − e⁻¹ を越える。「1 回引くとまれ」と「1 回生成するとまれ」は別の主張であり、利用者が出会うのは後者のほうだ。
平均してどれだけ驚くか
エントロピーは平均の驚き。交差エントロピーは、間違ったものに驚いたぶんの請求書だ。
確実なことが起きたと知っても何も得られないが、まず起きないことが起きたと知れば得るものは大きい。−log p はその考えに、独立な驚きが足し合わせられる算術を付けたものだ。確率を軸に沿ってドラッグしてほしい:
注目してほしい。p = 1 で驚きは 0、0.002 では 6.215 になり、 p が 0 へ向かうにつれて限りなく登っていく。0.693 の破線は 1 bit —— 公平なコインの驚き —— であり、同時に単位換算でもある。 nat に 1.4427 を掛ければ bit になる。
エントロピーは、その驚きを分布自身の重みで平均したものだ。各結果が −p log p だけ寄与する。その 6 つを端から端まで並べると、全長がそのままエントロピーになる。分布を集中させて、帯が縮むのを見てほしい:
一様が最大だ。等確率な 6 つの結果はH = 1.79 nat、ちょうどlog 6 を与え、帯は破線に届く。集中度をまで押すと p(1) は 0.831 まで登り、H は 0.73 まで落ちる。ほぼ心を決めた分布は、伝えるのにほとんど費用がかからない。
次はモデルだ。モデルは p を知らない。提案するのはq で、 p から届き続ける結果に対して −log q を払う。モデルを真から引き離して、帯の末尾に2 つ目のブロックが現れるのを見てほしい:
なぜなら q が p と同じだからで、ずれが 0 のとき 2 つの行は重なり、帯はちょうどH(p) = 1.67 になる。モデルは不確実性そのものの費用しか払っていない。 まで押すと請求は 2.05、うち 0.38 がずれだ。この超過が KL ダイバージェンスで、負にならず、 q が p のときだけ 0 になる —— 訓練目標として正しい理由がこれだ。
nat は感覚がつかみにくいので実務では exp(H) を報告する。「等確率な選択肢が何個ならこの難しさか」という数だ。マーカーは実在のモデルの損失から始まる:
損失 0 はパープレキシティ 1、つまり選ぶ余地がない。 GPT-2 の 50,257 トークンの語彙を当てずっぽうで引くと、破線の高さだ。 15 億パラメータの GPT-2 は WikiText-103 のゼロショットで 17.5 —— 2.86 nat、1 トークンあたり 4.13 bit —— なので、 5 万個の選択肢を 18 個ほどまで絞り込んだことになる。交差エントロピーとパープレキシティは、同じ 1 回の測定を 2 つの単位で言ったものだ。
ぜんぶを 1 枚に
上のどの節も、言語モデルが走らせる最後の 4 行に顔を出す。
順伝播の終わりはまさにこれだ。logits の並び、その softmax、 1 回の抽出、そして損失。ここでの真の次トークンは cat、サンプラーには固定の u = 0.62 が渡してある。そのトークンの logit をドラッグして、 4 つの段が一緒に動くのを見てほしい:
初期 logit 0.5 では、モデルが真のトークンに与える確率は 0.129 で、抽出は別のトークンを返す —— 損失は 2.05 nat、パープレキシティ 7.8 だ。logit を まで上げると確率は 0.973、抽出は cat に着地し、損失は 0.03 まで潰れる。 まで下げれば損失は 6.41 になる。
p = softmax(z / T) # z -= z.max() first i = searchsorted(cumsum(p), uniform()) loss = -log(p[true]) # nats, 0 at p = 1 ppl = exp(mean(loss)) # effective choices
訓練はその logit を動かし、サンプリングはその行を読む。この 4 行を 12 回まわせば、損失はただ足し合わさる —— §05 の「対数の和」に名前が付いただけだ。位置を 1 つずつ数え上げてほしい:
12 トークンぶんの合計は 19.70 nat、平均 1.641、 —— そのうち 3.91 は、p = 0.02 だった 7 番目の位置 1 つぶんだ。この 1 本の帯に主題のすべてが入っている。確率が記述し、対数がそれを足せるようにし、平均が論文の報告する数になる。