「はい/いいえ」の質問を重ねる
このコースの最初のレッスンで、人間がルールを書く「ルールベースAI」の例を見ました。
def judge(temperature):
if temperature >= 30:
return "暑い"
else:
return "普通"
決定木 は、このような if の分岐を、人間が書く代わりに データから自動で見つけ出す 分類モデルです。例えば迷惑メールの判定なら、学習の結果、次のような質問の組み合わせが見つかるかもしれません。
「無料」という単語を含む?
├─ はい → 差出人は連絡先に登録済み?
│ ├─ はい → 通常のメール
│ └─ いいえ → 迷惑メール
└─ いいえ → 通常のメール
木を逆さにしたような形をしているので「決定木」と呼ばれます。一番上の質問から始めて、答えに従って枝をたどり、一番下(葉)に書かれたクラスが予測結果になります。
では、コンピュータはどうやって「良い質問」を選ぶのでしょうか。直感的に言うと、質問で分けた後のそれぞれのグループが、なるべく1つのクラスだけになる(混ざり具合が減る) 質問を選びます。この「混ざり具合」を数値で表したものを 不純度 と呼び、scikit-learnでは標準で「ジニ不純度」という指標が使われます。
「良い質問」を計算で選ぶ:ジニ不純度
前のスライドで、決定木は「分けた後のグループの混ざり具合が減る質問」を選ぶと説明しました。この混ざり具合は、中学で習う分数と2乗の計算だけで求められます。
グループの中で、クラス0の割合を p0、クラス1の割合を p1 とすると、ジニ不純度 は次の式で計算します。
ジニ不純度=1−(p02+p12)
2つの例で確かめてみましょう。
- 3人全員が不合格のグループ: p0=1, p1=0 なので、1−(12+02)=0(まったく混ざっていない)
- 不合格2人・合格2人のグループ: p0=p1=21 なので、1−(41+41)=0.5(最も混ざっている)
つまり、ジニ不純度は 0に近いほど1つのクラスにまとまっていて、0.5に近いほど混ざっている ことを表します。
2つの質問を比べる
6人のデータ([勉強時間, 睡眠時間])で、次の2つの質問のどちらが良いかを計算してみます。
6人の勉強時間と睡眠時間の散布図。不合格の3人(○)は勉強時間1〜3時間、合格の3人(■)は6〜8時間にいて、勉強時間4.5時間の縦の線で完全に分かれます。睡眠時間では分かれていません。
グループごとの人数が違うので、人数の多いグループの結果ほど重く見る「重み付けした平均」で比べます。勉強時間の質問は0、睡眠時間の質問は0.5なので、混ざり具合をより減らせる勉強時間の質問 が選ばれます。
実際の決定木は、すべての項目について、隣り合う値のちょうど真ん中(今回なら3時間と6時間の間の4.5時間)を区切りの候補にし、すべての候補でこの計算を行って、一番小さくなる質問を選んでいます。コンピュータは、この単純な計算を大量に繰り返しているだけなのです。
scikit-learnで決定木を使う
決定木も、これまでと同じfit→predictの流れで使えます。モデルが変わっても使い方が共通していることが、scikit-learnの大きな利点です。
from sklearn.tree import DecisionTreeClassifier, export_text
import numpy as np
# [勉強時間, 睡眠時間] → 合格(1) / 不合格(0)
X = np.array([[1, 8], [2, 7], [3, 6], [6, 7], [7, 6], [8, 8]])
y = np.array([0, 0, 0, 1, 1, 1])
model = DecisionTreeClassifier(max_depth=2, random_state=0)
model.fit(X, y)
model.predict([[7, 7]]) # → array([1])
# 学習した分岐のルールを、人間が読める形で表示する
print(export_text(model, feature_names=["勉強時間", "睡眠時間"]))
export_textを実行すると、次のように学習したルールが表示されます。
|--- 勉強時間 <= 4.50
| |--- class: 0
|--- 勉強時間 > 4.50
| |--- class: 1
「勉強時間が4.5時間以下なら不合格、それより多ければ合格」というルールを、データから自動で見つけたことが分かります。このように、モデルが何を根拠に判断したかを人間が確認しやすい のが、決定木の大きな特徴です。
max_depthは、質問を重ねる段数の上限です。max_depth=2なら、最大2回の質問で答えを出します。今回のデータは1回の質問で完全に分けられたため、実際には1段の木になっています。random_stateは、結果を毎回同じにする(再現性を保つ)ための設定です。
よくある誤解:「深い木ほど賢い?」
質問をたくさん重ねられる深い木なら、どんなデータでも完璧に分類できそうに思えます。実際、max_depthを大きくすると、学習に使ったデータはほぼ100%正しく分類できるようになります。
しかし、これには落とし穴があります。深すぎる木は、学習データにたまたま含まれていた偶然のパターンや例外まで、細かいルールとして覚え込んでしまいます。その結果、新しいデータに対しては、かえって予測を外しやすくなる のです。この現象は「過学習」と呼ばれ、この章の最後のレッスンで詳しく扱います。
ロジスティック回帰との違い
どちらが常に優れているということはなく、データや目的によって使い分けます。
振り返り
- 決定木は、
ifの分岐にあたる質問をデータから自動で見つける
- 分けた後のグループのジニ不純度(混ざり具合)を人数で重み付けした平均が、最も小さくなる質問が選ばれる
max_depthで質問の段数を制限できる。深ければ良いわけではない