Table of Contents
決定木は、分類と回帰の両方のための最も直感的で広く使用されている機械学習アルゴリズムの一つです。 彼らは、機能値に基づいてブランチにデータを分割することによって働き、人間が決定を下す方法を模擬しています。 scikit-learnのようなライブラリは、建物の決定木を3重ねる一方で、最初から1つを実装することは、アルゴリズムの内部作業を把握するための優れた方法です。 このチュートリアルでは、理論とコードを説明します。そのため、地面から独自の決定木を組み立てることができます。
決定の木とは何ですか?
決定ツリーは、各内部ノードが機能に関するテストを表すフローチャートのような構造です(例えば、「Is age > 30?」)、各ブランチは、そのテストの結果を表し、各リーフノードはクラスラベルまたは連続値を保持しています。 目標は、データ機能から推論された簡単な決定ルールを学ぶことで、ターゲット変数を予測するモデルを作成することです。 意思決定ツリーは、データプリプロダリング(スケールまたは正規化)を解釈し、必要とされるのが簡単なため人気があります。
ツリーは再帰的に構築されます。ルートから始まり、アルゴリズムは最もクリーンなデータを分離する最高の機能と分岐点を選択します。このプロセスは、停止条件が満たされるまで、各サブセットに繰り返します。 背景が高まれば、Wikipediaの]の決定ツリー学習に関するエントリは、固体概要を提供します。
理解しなければならないコアコンセプト
ノデックス、ブランチ、および葉
ルートノードは、トレーニングデータセット全体が含まれています。内部ノードは、機能をテストし、データを2つ以上の子ノードに分割します。ブランチは、テストの結果を表す接続です。リーフノード(ターミナルノード)は、最終予測を出力します。分類の最も一般的なクラスまたは回帰中の平均値。
分裂の基準
ツリーを作成するには、潜在的な分割の品質を測定する方法が必要です。最も一般的な基準は次のとおりです。
- []Gini impurity] - 分類で、ランダムに分類された要素がサブセットのクラスの分布に応じてランダムにラベル付けされたかどうかを誤ってラベル付けされる。 より低いGiniはより良いです。
- [Entropy[]] - 障害や不確実性をセットで測定します。 目標は、分割後の不適切性(情報収集)を最小限に抑えることです。
- ] 回帰木に用いられる分散削減 。分割で得られる分散(または四角形エラー)の減少を計算します。
アルゴリズムは、あらゆる機能に分割可能な評価を行い、不純物(または情報を得る)の最大の削減をもたらすものを選ぶ。
情報ゲインとゲイン比率
情報ゲインは、親ノードの不純物と子供の不純物の重みのある合計の違いです。単純に、それは多くの値で機能に好む傾向があります。ゲイン比(C4.5で使用されます)はこれを正規化します。このチュートリアルでは、CART(分類および回帰ツリー)でデフォルトであるGini不純物を使用して標準的な情報ゲインに固執します。
段階的に意思決定ツリーのステップを造る
1. データの準備
機能とターゲットラベルのデータセットが必要です。 単純性のために、バイナリ分類データセットを数値機能で使用します。 たとえば:
- 特集: 年齢、収入
- 対象:)承認済み (1) または承認されていない (0)
データの消去: 不足している値を処理する、重複を削除し、数値型を確保します。 決定ツリーは、混合されたデータ型を処理することができますが、実装の数値に固執します。
2. 分裂のクトリタリオン機能を定義する
Gini の不純物を実装します。Gini のインデックスは、以下の項目です。
[]]
p i がクラス i の項目の割合である。バイナリの分割では、Gini 全体が子ノードの重み平均である。
3. 分割評価の実施
各機能に対して、独自の値をソートします。 可能なすべてのしきい値(連続ソート値の中間点)をテストします。 各候補のしきい値については、データを左右のグループに分割し、Giniを計算し、最良の分割を追跡します。
4. ツリーを復活させる
データのサブセットと現在の深さを取る関数を作成します。 ストップ条件(例、最大深さが到達した、最小サンプル/ノードまたは情報ゲインなし)を確認します。 条件が満たされた場合、大部分のクラスでリーフノードを作成します。 それ以外の場合は、最高の分割を見つけて内部ノードを作成してから、左と右に関数を呼び出します。
5. 予測をする
ツリーが構築されると、予測は簡単です。ルートで始まり、新しいサンプルで機能テストを評価し、あなたが着陸する葉の値を返します。
Python のフル実装
以下は、Gini不純物を使用して分類するための決定ツリーの完全な最小限の実装です。このコードは学習のために意味されています。それは大きなデータセットのために最適化されていません。
import numpy as np
from collections import Counter
class DecisionTree:
def __init__(self, max_depth=None, min_samples_split=2):
self.max_depth = max_depth
self.min_samples_split = min_samples_split
self.tree = None
def fit(self, X, y):
dataset = np.column_stack((X, y))
self.tree = self._grow_tree(dataset)
def _grow_tree(self, dataset, depth=0):
X, y = dataset[:, :-1], dataset[:, -1]
n_samples, n_features = X.shape
n_labels = len(np.unique(y))
# Stopping conditions
if (n_labels == 1 or depth == self.max_depth or n_samples < self.min_samples_split):
leaf_value = Counter(y).most_common(1)[0][0]
return {'leaf': True, 'value': leaf_value}
best_feature, best_threshold = self._best_split(dataset, n_features)
if best_feature is None:
leaf_value = Counter(y).most_common(1)[0][0]
return {'leaf': True, 'value': leaf_value}
left_idx, right_idx = self._split(dataset[:, best_feature], best_threshold)
left_subtree = self._grow_tree(dataset[left_idx], depth+1)
right_subtree = self._grow_tree(dataset[right_idx], depth+1)
return {'leaf': False,
'feature': best_feature,
'threshold': best_threshold,
'left': left_subtree,
'right': right_subtree}
def _best_split(self, dataset, n_features):
best_gini = float('inf')
best_feature, best_threshold = None, None
for feature in range(n_features):
thresholds = np.unique(dataset[:, feature])
for i in range(len(thresholds)-1):
thresh = (thresholds[i] + thresholds[i+1]) / 2
left_idx, right_idx = self._split(dataset[:, feature], thresh)
if len(left_idx) == 0 or len(right_idx) == 0:
continue
gini = self._gini_gain(dataset, left_idx, right_idx)
if gini < best_gini:
best_gini = gini
best_feature = feature
best_threshold = thresh
return best_feature, best_threshold
def _split(self, values, threshold):
left_idx = np.where(values <= threshold)[0]
right_idx = np.where(values > threshold)[0]
return left_idx, right_idx
def _gini_gain(self, dataset, left_idx, right_idx):
total = len(left_idx) + len(right_idx)
gini_left = self._gini(dataset[left_idx, -1])
gini_right = self._gini(dataset[right_idx, -1])
return (len(left_idx)/total) * gini_left + (len(right_idx)/total) * gini_right
def _gini(self, labels):
_, counts = np.unique(labels, return_counts=True)
p = counts / np.sum(counts)
return 1 - np.sum(p**2)
def predict(self, X):
return np.array([self._predict_row(x, self.tree) for x in X])
def _predict_row(self, x, node):
if node['leaf']:
return node['value']
if x[node['feature']] <= node['threshold']:
return self._predict_row(x, node['left'])
else:
return self._predict_row(x, node['right'])
ツリーをテストする
古典的なアイリスのデータセット(バイナリ分類の2つの機能)のように単純なデータセットを使用してください。 []]scikit-learn Iris dataset[がうまく機能します。 ツリーの精度をscikit-learnのと比較すると、正しい精度が確認されます。
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
data = load_iris()
X = data.data[:100] # take only first two classes (binary)
y = data.target[:100]
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2)
tree = DecisionTree(max_depth=3)
tree.fit(X_train, y_train)
preds = tree.predict(X_test)
accuracy = np.mean(preds == y_test)
print(f'Accuracy: {accuracy:.2f}')
あなたの木を改善する高度なテクニック
過度の苦しみを避けるために剪定
完全に成長したツリーは、トレーニングデータでノイズを記憶することができます。剪定は、予測力が少ない枝を取り除きます。一般的な方法は、事前に剪定(または)、およびポストプルーン(検証セットまたは費用対複雑化剪定を使用して、完全なツリーを除去する)です。当社の実装は、すでに事前剪定をサポートしています。
連続的かつ定性的特徴の取り扱い
連続機能では、ソートされた値と値の境界線を区別します。分類機能(例えば「色=赤/緑/青」)では、各カテゴリは別々のブランチ(マルチ・ウェイ・スプリット)になるか、バイナリ・エンコードをすることができます。ほとんどの近代的な実装(scikit-learnのような)は、すべてのサブセットを評価することによって、分類機能でもバイナリスプリットを使用します。
価値を逃すことで対処
実際のデータでは、多くの場合、値が不足しています。シンプルなアプローチは、機能を持つトレーニングサンプルの中で最も頻繁にブランチに欠落した値を指定することです。C4.5は確率的メソッドを使用します。これは初心者のチュートリアルなので、データが完成していると仮定します。
図書館とさらに読むことと比較
ゼロから構築することは、教育的である一方で、生産システムは最適化されたCの実装を提供するscikit-learnなどのライブラリを使用します。あなたは公式[からもっと学ぶことができます。 scikit-learn決定の木文書]。 より深い理論のために、ハシー、チブスラニ、フリードマンによる「統計学習の要素」は、権威あるリソースです。 もう一つの優れた言及は、Breti ARTによって元のCブックです。
コンテンツ
ゼロから決定ツリーを構築することは、機械学習における最も基本的なアルゴリズムの1つです。 あなたは、単純な再帰的な分割手順が強力なモデルを作り出すことができる方法を学びました。 自分でコードを書くことによって、あなたは衝動の対策、分割選択、バイアスと分散間の取引オフのより深い理解を得ることができます。 次のステップとして、回帰サポート、剪定、または分類機能の取り扱いを追加します。 あなたがここに開発するスキルは、あなたがより複雑な方法を上回るのと同様に役立つでしょう。