—
type: note
title: 决策树编程作业
date: 2026-08-27
description: C2W4 从零造一棵决策树认蘑菇:熵量"有多乱"、分裂掰两拨、信息增益=切前熵−切后加权熵、选增益最大特征。坑:0log₂(0) 要 else 兜底(np.log2(0)→-inf→nan);get_best_split 占位数不能撞真实下标。
—
# 2026-08-27 学习笔记
## 笔记区(学的时候随手记)
### 核心观点
**造一棵决策树,让它自己学会认蘑菇。** 数据:10 朵蘑菇、3 个特征(棕色菌盖 / 细长菌柄 / 独生),标签是**可食(1) / 有毒(0)**。这棵树的活儿就是:**每次问一个问题,把蘑菇切得更干净,切到分不动为止。**
造树不靠背,靠 4 块积木(作业就是补这 4 个函数):
| 积木 | 干什么 | 人话 |
| ———————————– | ——————– | ———————- |
| 熵 `compute_entropy` | 量这一拨"有多乱" | 一团糊 vs 分得清 |
| 分裂 `split_dataset` | 按某特征掰成左右两拨 | 一个问题把蘑菇分开 |
| 信息增益 `compute_information_gain` | 这一刀划不划算 | 切完比切前干净多少 |
| 选最佳 `get_best_split` | 挑增益最大的特征 | 试遍问题,挑最能分开的 |
### 关键方法/流程
– **熵**:全一种 = 0,一半一半 = 1。公式 `-p1·log₂(p1) – (1-p1)·log₂(1-p1)`,`p1` = 可食比例。
– **信息增益** = 分裂前熵 − 分裂后两拨加权熵。**越大 = 混乱减少越多 = 越该往下分。**
– **加权**:两拨大小可能不同,要按占比加权,不能直接平均熵。
### 实战记录
**① 熵里 `0·log₂(0)=0` 要 else 兜底**
> 0log₂(0)=0 是人为约定。但计算机里 np.log2(0) 返回 -inf,0 × -inf = nan,所以 p1=0 或 p1=1 时公式会算出 nan 而不是 0。因此必须写 `else: entropy = 0` 手动定死,否则过不了测试。
**② `best_feature=-1` 占位不撞真实下标**
> 在 get_best_split 中用 -1 做占位,因为 -1 不是任何真实特征下标(本数据集下标是 0、1、2),能安全表示"还没选到"。若用 0 当初始值,0 本身是合法特征下标,一旦循环没找到更大增益,会误返回特征 0。关键是占位数别撞上真实下标——-1、3、4、5 都可以,**2 不行**(2 是"独生"特征的下标)。
## 代码逐段展开
### Ex1 `compute_entropy(y)` — 算一拨的熵
```python
def compute_entropy(y):
entropy = 0.
if len(y) != 0: # 空节点返回 0
p1 = len(y[y == 1]) / len(y) # 可食比例
if p1 != 0 and p1 != 1:
entropy = -p1 * np.log2(p1) – (1-p1) * np.log2(1-p1)
else: # p1=0/1 → 0log₂(0) 约定
entropy = 0.
return entropy
```
– `p1 = len(y[y == 1]) / len(y)`:数可食的 ÷ 总数 = 比例。
– `p1` 是 0 或 1 时 `np.log2` 会算出 `-inf`,`0×-inf = nan`,必须 `else` 兜底给 0。
– 根节点 5 可食 5 有毒 → p1=0.5 → 熵 = 1.0(最乱),测试预期就是 1.0。
### Ex2 `split_dataset(X, node_indices, feature)` — 按特征掰两拨
```python
def split_dataset(X, node_indices, feature):
left_indices = []
right_indices = []
for i in node_indices: # 逐个样本
if X[i][feature] == 1: # 特征=1 左
left_indices.append(i)
else: # 特征=0 右
right_indices.append(i)
return left_indices, right_indices
```
– 只记**下标**,不搬数据。按"棕色菌盖"(feature=0)分根节点:左 `[0,1,2,3,4,7,9]`,右 `[5,6,8]`。
### Ex3 `compute_information_gain(…)` — 这刀划不划算
```python
def compute_information_gain(X, y, node_indices, feature):
left_indices, right_indices = split_dataset(X, node_indices, feature)
X_node, y_node = X[node_indices], y[node_indices]
X_left, y_left = X[left_indices], y[left_indices]
X_right, y_right = X[right_indices], y[right_indices]
information_gain = 0
node_entropy = compute_entropy(y_node) # 分裂前多乱
left_entropy = compute_entropy(y_left) # 左拨
right_entropy = compute_entropy(y_right) # 右拨
w_left = len(X_left) / len(X_node) # 左占比
w_right = len(X_right) / len(X_node) # 右占比
weighted_entropy = w_left * left_entropy + w_right * right_entropy
information_gain = node_entropy – weighted_entropy
return information_gain
```
– 分裂用的是**你自己写的 Ex2**,熵用**你自己写的 Ex1**——积木互相拼。
– 两拨大小不一样 → 按占比加权,不能直接平均。
– 预期:棕色 0.0349、细长 0.1245、独生 0.2781 → **独生收益最大**。
### Ex4 `get_best_split(…)` — 挑增益最大的特征
```python
def get_best_split(X, y, node_indices):
num_features = X.shape[1]
best_feature = -1
max_info_gain = 0
for feature in range(num_features): # 试遍每个特征
info_gain = compute_information_gain(X, y, node_indices, feature)
if info_gain > max_info_gain:
max_info_gain = info_gain
best_feature = feature
return best_feature
```
– `best_feature = -1`:**占位**,不撞真实下标,安全表示"还没选到"。
– 循环里 `if info_gain > max_info_gain`:找到更大的就换,最后留最大的。
– 根节点:独生(2) 增益 0.2781 最大 → 第一次分裂就选"独生"。
## 线索区(学完后合上材料,自问自答)
> 先看问题 → 自己回答 → 对照。答不上来的就是没学透。
**Q: 熵最大是啥时候?为什么?**
A: p1=0.5(一半一半)时熵=1 最大。最乱就是两种各占一半,猜哪边都没把握;全是一种则熵=0,一点悬念没有。
**Q: 信息增益越大说明什么?**
A: 这一刀把混乱削减得越多,越值得在这里往下分。本质 = 分裂前熵 − 分裂后加权熵。
**Q: 占位数为什么不能用 0 或 2?**
A: 0、2 都是真实特征下标。占位数若撞真实下标,循环里一直没找到更大增益时,会把"还没选到"误当成"选到了那个特征"。用 -1(或 3、4、5)才安全。
**Q: 信息增益里为什么按人数加权,而不是两个熵直接平均?**
A: 左拨右拨大小可能不一样。直接平均=把 10 人和 1 人当等重,不公平。按占比加权才反映"数据多的那边说了算"。
## 总结(50 字以内)
C2W4 决策树 = 从零造树:熵量多乱、分裂掰两拨、信息增益=切前熵−切后加权熵、选增益最大特征。坑:0log₂(0) 要 else 兜底、占位数不撞真实下标。
—
## 复习卡片
| 概念 | 一句话 | 我的场景 |
| ——– | ———————————- | ————————— |
| 熵 | 量这一拨多乱;全一种=0、一半一半=1 | 判断该不该继续切 |
| 信息增益 | 切前熵 − 切后加权熵 | 越大越该往下分 |
| 加权熵 | 按占比加权,不直接平均 | 两拨人数不一样时 |
| 占位数 | -1/3/4/5 行,0、2 不行 | get_best_split 的"还没选到" |
| 0log₂(0) | 数学约定=0,机器算 nan | np.log2(0) 要 else 兜底 |
网硕互联帮助中心




评论前必须登录!
注册