10個の点数を、確率に変える
CNNの最後の層は、0〜9の各クラスについて1つずつ、合わせて10個の 点数 を出します。点数が大きいクラスほど、「その数字らしい」という意味です。ただし、点数は −5.2 や 5.66 のような数で、このままでは「何%くらいその数字らしいか」が分かりません。
第5章のロジスティック回帰では、2クラスの分類で、点数 z をシグモイド関数で0〜1の確率に変えました。10クラスのときに使うのが、ソフトマックス関数 です。
ソフトマックス関数の手順
- 10個の点数 z それぞれについて、ez を計算する(e は第5章で使った約2.718という決まった数)
- 10個の ez を、全部足す
- それぞれの ez を、その合計で割る
式で書くと、クラス i の確率は次のようになります。
pi=ez0+ez1+⋯+ez9ezi
3クラスで計算してみる
点数が (2,1,0) の3クラスで計算します。第5章の表のとおり、e2≈7.39、e1≈2.72、e0=1 です。
- 合計: 7.39+2.72+1=11.11
- クラス0の確率: 7.39÷11.11≈0.665
- クラス1の確率: 2.72÷11.11≈0.245
- クラス2の確率: 1÷11.11≈0.090
点数が 2, 1, 0 の3クラスにソフトマックス関数を使った結果。クラス0が0.665、クラス1が0.245、クラス2が0.090
3つの確率を足すと、次のように1になります。全部の割合を合わせると1(100%)になるように、合計で割っているからです。
0.665+0.245+0.090=1
なぜ ez を使うのか
- ez は、z が負の数でも必ず正の数になります。確率が負になることはありません
- z が大きいほど ez も大きくなるので、点数の順番がそのまま確率の順番になります
- z が1増えるごとに約2.7倍になるので、点数の差が少し開くだけで、確率の差ははっきり開きます
クラスが2つで、片方の点数を0と決めたとき、もう片方の確率は ez+e0ez=1+e−z1(分母と分子を ez で割った形)となり、第5章のシグモイド関数とまったく同じになります。シグモイド関数を、たくさんのクラスに広げたものと考えられます。
実際のCNNの出力と、学習の罰
手書き数字のデータで学習させたCNN(フィルター16枚)に、テストデータの画像を1枚入力しました。正解は数字の8です。最後の層が出した10個の点数と、ソフトマックス関数で確率に変えた結果は次のとおりです(点数は小数第2位まで、確率は小数第3位まで。どちらも、その次の位を四捨五入)。
正解が8の画像に対するCNNの確率。8が0.952と飛び抜けて高く、次に高い4でも0.028。ほかのクラスはどれも0.01未満
CNNは「95.2%の確率で8」と答えています。予測としては、一番確率が高いクラス(8)を選びます。2番目に高いのは4(2.8%)でした。
学習では、正解のクラスの確率を見る
学習の罰(損失)は、第5章と同じ 交差エントロピー です。第5章で学んだとおり、「正解に付けた確率が低いほど大きくなる罰」で、正解に付けた確率が半分になるたびに、罰が約0.69ずつ増えます。
10クラスあっても、罰の計算に使うのは 正解のクラスの確率だけ です。ただし、確率は合計が1なので、正解の確率を上げると、自然にほかのクラスの確率は下がります。
学習では、訓練データ全体の罰の平均が小さくなるように、重みを少しずつ動かします。第4章の勾配降下法、第6章の誤差逆伝播と、考え方はまったく同じです。
CNNの学習も、誤差逆伝播で行う
第6章で、誤差逆伝播は「変化の割合を、後ろの層から掛け算でつないでいく方法」だと学びました。CNNでも同じです。ここでは、後ろから順に、各部分で何が起きるかを言葉で追いかけます(細かい式は、次のレッスンのコードで動かして確かめます)。
- ソフトマックスと交差エントロピー: 罰の変化の割合は、とても簡単な形になります。各クラスについて、予測した確率から、正解のクラスなら1を、それ以外のクラスなら0を引いた値です。前のスライドの例なら、正解の8は 0.952−1=−0.048、4は 0.028−0=0.028 です。「正解の確率をもっと上げ、ほかを下げる向き」を表しています
- 最後の層: 第6章の層とまったく同じです
- プーリング: 最大値プーリングでは、2×2のまとまりの中で 一番大きかったマスにだけ、変化の割合を戻します。ほかの3マスは、出力に影響していなかったので0です
- ReLU: 第6章と同じで、0より大きかったマスだけ、変化の割合をそのまま通します
- 畳み込み: フィルターの1つの重みは、窓の すべての位置 で使われていました(重みの共有)。そのため、その重みの変化の割合は、すべての位置の分を足し合わせたもの になります
このうち、第6章と違うのは3と5だけで、どちらもこの章で学んだ仕組み(一番大きいマスだけを残す、同じ重みを使い回す)から、自然に決まります。
重みの共有があっても、学習の考え方は変わらない
5の「すべての位置の分を足し合わせる」は、難しいことではありません。第4章で、損失は「全部のデータの罰の合計(平均)」でした。1つの重みが何か所で使われていても、「その重みを少し動かしたら、全部の場所の罰の合計がどれだけ変わるか」を、場所ごとに計算して足せばよいのです。
実際には、ライブラリが自動で計算する
この計算を手で書くのは大変ですが、PyTorchなどのライブラリは、順伝播の計算をたどって、変化の割合を自動で計算してくれます(自動微分 と呼ばれます)。私たちが書くのは、順伝播(モデルの形)と、損失の計算だけです。次のスライドで、この章で学んだものがPyTorchのどの書き方にあたるかを確かめます。
PyTorchとの対応と、振り返り
この章で自分の手で計算したものは、実務でよく使われるライブラリ PyTorch では、次のように書きます。このコースの演習の実行環境(ブラウザで動くPython)ではPyTorchを動かせないため、ここでは書き方の対応だけを確かめます。
組み立てると、この章のCNNは次のように書けます。
import torch
from torch import nn
model = nn.Sequential(
nn.Conv2d(1, 16, kernel_size=3), # 8×8 → 6×6 が16枚
nn.ReLU(),
nn.MaxPool2d(2), # 6×6 → 3×3
nn.Flatten(), # 3×3×16 = 144個に並べる
nn.Linear(144, 10), # 10クラスの点数
)
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
# 学習の1回分(images は画像、labels は正解)
scores = model(images) # 順伝播
loss = loss_fn(scores, labels) # ソフトマックス+交差エントロピー
optimizer.zero_grad() # 前回の変化の割合を消す
loss.backward() # 誤差逆伝播
optimizer.step() # 重みを動かす
このモデルの重みとバイアスの数は、レッスン5で数えたとおり1,610個です。
よくある誤解
- 「
nn.CrossEntropyLoss() を使うときは、モデルの最後にソフトマックスを入れる」: nn.CrossEntropyLoss() は、中でソフトマックスも計算します。モデルの最後にもソフトマックスを入れると、2回計算することになり、学習がうまく進みません。モデルは点数を出すところまでにします
- 「ソフトマックスの確率は、その答えが正しい確率そのもの」: 0.952は「このモデルがどれくらい自信を持っているか」を表す数で、本当に95.2%の割合で正解するとは限りません。学習データにない種類の画像にも、高い確率を付けてしまうことがあります
- 「ライブラリを使えば、中の仕組みは知らなくてよい」: 仕組みを知っていると、形が合わないエラーや、学習が進まない原因を見つけられます。この章で手で計算したことは、そのための土台です
振り返り
- ソフトマックス関数は、点数を「合計が1になる確率」に変える(ez を合計で割る)。シグモイド関数を、たくさんのクラスに広げたもの
- 学習の罰は、正解のクラスの確率を見る交差エントロピー
- CNNの学習も誤差逆伝播。プーリングは一番大きいマスにだけ、共有した重みは全部の位置の分を足して、変化の割合を戻す