想像你在玩一個「20 個問題」的遊戲——透過一連串「是/否」問題縮小範圍,最終猜出答案。決策樹正是用同樣的邏輯來做預測:它把預測空間反覆切成小方塊,每切一刀就是一個問題(例如「年資 < 4.5 年?」),最終落在同一區塊的資料就給出相同的預測值。
課本以 Hitters 資料集為例:用 Years(大聯盟年資)和 Hits(去年安打數)來預測球員薪資。一棵簡單的回歸樹長這樣:
Years < 4.5
/ \
$165,174 Hits < 117.5
/ \
$402,834 $845,346
這棵樹把球員分成三個區域:
建立回歸樹的目標是:把預測空間切成 J 個不重疊的矩形區域 R₁, R₂, …, Rⱼ,使得每個區域內的觀測值盡可能相似。數學上就是最小化 RSS(殘差平方和):
其中 \(\hat{y}_{R_j}\) 是第 j 個區域內所有訓練觀測值的 平均值。
但問題來了:要考慮所有可能的切割方式,計算量太大。因此採用遞迴二元分割——一個貪婪(greedy)策略:
對於每個預測變數 Xⱼ 和每個可能的切割點 s,定義兩個半平面:
然後選擇使兩區域 RSS 之和最小的 (j, s):
重複此過程:每次選擇一個現有區域,一刀切下去,直到滿足停止條件(如每個區域少於 5 筆資料)。
如果讓樹一直長下去,它會完美記憶訓練資料——也就是過度擬合。解法是:先讓樹長到最大(T₀),再用「成本複雜度修剪」(cost complexity pruning)把它修小。
這裡 |T| 是終端節點數,α 是調節參數:
分類樹的結構和回歸樹一樣,差別在於:
| 指標 | 公式 | 特性 |
|---|---|---|
| 分類錯誤率 | \(E = 1 - \max_k(\hat{p}_{mk})\) | 直觀但對樹的生長不夠敏感;適合用於修剪後的評估 |
| Gini 指數 | \(G = \sum_{k=1}^{K} \hat{p}_{mk}(1 - \hat{p}_{mk})\) | 衡量總變異;值越小節點越純。sklearn 預設使用 |
| 熵(Entropy) | \(D = -\sum_{k=1}^{K} \hat{p}_{mk} \log \hat{p}_{mk}\) | 資訊理論概念;數值上與 Gini 很接近 |
課本以 Heart 資料集為例(303 位胸痛病人,預測是否患有心臟病),交叉驗證選出一個六葉節點的分類樹。樹中的切割變數包括 Thal(鉈壓力測試)、Ca(鈣含量)、MaxHR(最大心率)等。
try:
from google.colab import drive
drive.mount('/content/drive')
DATA_PATH = '/content/drive/MyDrive/ISLP_data/'
except ImportError:
DATA_PATH = '/tmp/'
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from sklearn.tree import DecisionTreeRegressor, plot_tree
from ISLP import load_data
# 載入 Hitters 資料
hitters = load_data('Hitters')
hitters = hitters.dropna(subset=['Salary'])
hitters['logSalary'] = np.log(hitters['Salary'])
# 用 Years 和 Hits 建立回歸樹(對應課本圖 8.1)
X = hitters[['Years', 'Hits']].values
y = hitters['logSalary'].values
# 限制最大葉節點數為 3(對應課本的 3 區域樹)
tree3 = DecisionTreeRegressor(max_leaf_nodes=3, random_state=42)
tree3.fit(X, y)
print(f"R^2 (training): {tree3.score(X, y):.3f}")
print(f"樹的深度: {tree3.get_depth()}, 葉節點數: {tree3.get_n_leaves()}")
# 可視化
plt.figure(figsize=(10, 5))
plot_tree(tree3, feature_names=['Years', 'Hits'], filled=True,
rounded=True, fontsize=9)
plt.title('Hitters 回歸樹 (max_leaf_nodes=3)')
plt.tight_layout()
plt.savefig('/tmp/hitters_regression_tree.png', dpi=100)
plt.show()
try:
from google.colab import drive
drive.mount('/content/drive')
DATA_PATH = '/content/drive/MyDrive/ISLP_data/'
except ImportError:
DATA_PATH = '/tmp/'
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import numpy as np
from sklearn.tree import DecisionTreeRegressor
from sklearn.model_selection import cross_val_score, train_test_split
from ISLP import load_data
# 載入 Hitters,使用全部數值特徵
hitters = load_data('Hitters')
hitters = hitters.dropna(subset=['Salary'])
hitters['logSalary'] = np.log(hitters['Salary'])
features = ['AtBat', 'Hits', 'HmRun', 'Runs', 'RBI',
'Walks', 'Years', 'CAtBat', 'CHits', 'CHmRun',
'CRuns', 'CRBI', 'CWalks', 'PutOuts', 'Assists', 'Errors']
X = hitters[features].values
y = hitters['logSalary'].values
X_train, X_test, y_train, y_test = train_test_split(
X, y, test_size=0.5, random_state=1)
# 長出大樹
big_tree = DecisionTreeRegressor(random_state=42)
big_tree.fit(X_train, y_train)
print(f"大樹葉節點數: {big_tree.get_n_leaves()}")
# 取得 cost complexity pruning path
path = big_tree.cost_complexity_pruning_path(X_train, y_train)
ccp_alphas = path.ccp_alphas[:-1]
# 用 cross-validation 找最佳 alpha
cv_scores = []
for alpha in ccp_alphas:
tree = DecisionTreeRegressor(ccp_alpha=alpha, random_state=42)
scores = cross_val_score(tree, X_train, y_train, cv=5,
scoring='neg_mean_squared_error')
cv_scores.append(-scores.mean())
best_alpha = ccp_alphas[np.argmin(cv_scores)]
print(f"最佳 alpha: {best_alpha:.6f}")
# 用最佳 alpha 訓練最終樹
final_tree = DecisionTreeRegressor(ccp_alpha=best_alpha, random_state=42)
final_tree.fit(X_train, y_train)
print(f"修剪後葉節點數: {final_tree.get_n_leaves()}")
test_mse = np.mean((y_test - final_tree.predict(X_test))**2)
print(f"Test MSE: {test_mse:.4f}")
# 繪圖
plt.figure(figsize=(10, 5))
plt.plot(ccp_alphas, cv_scores, 'b-o', markersize=4, label='CV MSE')
vline = plt.axvline(best_alpha, color='red', linestyle='--',
label=f'Best alpha={best_alpha:.4f}')
plt.xlabel('alpha (cost complexity parameter)')
plt.ylabel('Cross-validated MSE')
plt.title('Cost Complexity Pruning: CV Error vs alpha')
plt.legend()
plt.grid(True, alpha=0.3)
plt.tight_layout()
plt.savefig('/tmp/pruning_cv.png', dpi=100)
plt.show()
try:
from google.colab import drive
drive.mount('/content/drive')
DATA_PATH = '/content/drive/MyDrive/ISLP_data/'
except ImportError:
DATA_PATH = '/tmp/'
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from sklearn.tree import DecisionTreeClassifier, plot_tree
from sklearn.model_selection import cross_val_score, GridSearchCV
from ISLP import load_data
# 載入 Heart 資料(從 statlearning.com 下載的 CSV)
heart = pd.read_csv(DATA_PATH + 'Heart.csv')
# 準備特徵和目標
X_df = pd.get_dummies(heart.drop('AHD', axis=1), drop_first=True)
X = X_df.values
y = (heart['AHD'] == 'Yes').astype(int).values
# 用 CV 選擇最佳葉節點數
param_grid = {'max_leaf_nodes': range(2, 21)}
grid = GridSearchCV(
DecisionTreeClassifier(random_state=42),
param_grid, cv=5, scoring='accuracy'
)
grid.fit(X, y)
print(f"最佳葉節點數: {grid.best_params_['max_leaf_nodes']}")
print(f"CV 準確率: {grid.best_score_:.3f}")
print(f"訓練準確率: {grid.score(X, y):.3f}")
# 可視化最佳分類樹
best_tree = DecisionTreeClassifier(
max_leaf_nodes=grid.best_params_['max_leaf_nodes'],
random_state=42
)
best_tree.fit(X, y)
plt.figure(figsize=(14, 7))
plot_tree(best_tree, feature_names=X_df.columns,
class_names=['No HD', 'HD'], filled=True, rounded=True, fontsize=7)
plt.title(f'Heart 分類樹 (max_leaf_nodes={grid.best_params_["max_leaf_nodes"]})')
plt.tight_layout()
plt.savefig('/tmp/heart_classification_tree.png', dpi=100)
plt.show()
# 比較 Gini vs Entropy
for criterion in ['gini', 'entropy']:
clf = DecisionTreeClassifier(criterion=criterion, random_state=42)
scores = cross_val_score(clf, X, y, cv=5, scoring='accuracy')
print(f"{criterion}: CV accuracy = {scores.mean():.3f} "
f"(+/- {scores.std():.3f})")
決策樹是最受醫生歡迎的 ML 模型——因為它可以畫出來。一棵心臟病風險決策樹可以印在一張紙上,醫生按圖索驥:先看 thal 測試結果 → 再看最大心率 → 再看膽固醇...,最終得出風險評估。這種透明性在醫療領域至關重要。
銀行審核貸款時,決策樹可以提供明確的審核路徑:「收入 < 3 萬 → 負債比 > 40% → 拒絕」。比起黑箱深度學習模型,決策樹可以給客戶一個清楚的拒絕理由,符合金融監管對「可解釋性」的要求。
行銷團隊用決策樹將客戶分群:「年消費 > 5000 且最近購買天數 < 30 → VIP 客戶」。行銷人員不需要懂統計,看樹狀圖就能理解客戶分類邏輯。
| 方法 | 預測方式 | 可解釋性 | 非線性 | 典型場景 |
|---|---|---|---|---|
| 線性迴歸 | 全域線性函數 | ⭐⭐⭐⭐⭐ | ❌ 需手動加多項式 | 有理論支持的線性關係 |
| KNN | 局部鄰近平均 | ⭐⭐ | ✅ 自然非線性 | 樣本少、邊界不規則 |
| 決策樹 | 分段常數 | ⭐⭐⭐⭐ | ✅ 階梯狀逼近 | 需要解釋的場景 |
| GAM(廣義加法模型) | 平滑函數之和 | ⭐⭐⭐⭐ | ✅ 平滑非線性 | 可加性成立的複雜關係 |
| SVM with RBF kernel | 最大邊界超平面 | ⭐ | ✅ 高度非線性 | 高準確度要求 |