You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

ID3算法中的递归机制及相关熵计算Python代码解析请求

解析ID3算法的递归机制与样本标签概率、熵计算代码解读

一、ID3算法的递归机制拆解

ID3构建决策树的核心就是递归拆分样本集,说白了就是一步步把“杂乱的大样本堆”拆成“整齐的小样本堆”,直到没法拆或者样本都统一了为止。核心逻辑可以拆成两部分:

递归的终止信号(停拆条件)

  • 当当前手里的所有样本标签完全一致时,直接把这个标签当叶子节点,不用再拆分了
  • 要么就是所有能用的特征都用完了,或者样本集为空,那就选当前样本里占比最高的标签当叶子节点

递归的拆分逻辑(拆树步骤)

  1. 先算当前样本集的信息熵——这玩意儿就是衡量样本“乱不乱”的指标,熵越高说明样本标签越混杂
  2. 挨个试剩下的没用来拆分的特征,计算每个特征拆分后能让熵减少多少(也就是信息增益)
  3. 挑信息增益最大的那个特征当当前节点的“拆分依据”
  4. 按这个特征的不同取值,把样本分成好几堆子样本
  5. 对每一堆子样本,重复上面的步骤,递归构建子树

简单说就是“选最能让样本变整齐的特征→拆成小堆→给每堆重复操作”,直到触发终止条件。

二、样本标签概率与熵计算的Python代码解读

我逐行给你掰扯这段代码的逻辑,顺便指出里面的小问题:

1. 训练集说明

import numpy as np
training_set = np.array([[0, 1, 0, 1, 0, 1],[0, 0, 0, 1, 0, 0],[0, 0, 0, 0, 1, 0],[1, 0, 1, 0, 1, 0],[0, 1, 1, 1, 0, 1],[0, 1, 0, 0, 1, 1],[1, 1, 1, 0, 0, 0],[1, 1, 1, 1, 0, 1],[0, 1, 1, 0, 1, 0],[1, 1, 0, 0, 0, 1],[1, 0, 0, 0, 1, 0]])

这是一个11行6列的numpy数组,每一行是一个样本,最后一列是样本的标签(0或1),前5列是样本的特征值,总共包含11个训练样本。

2. p(X)函数:计算标签的概率分布

这个函数的作用是统计样本集中标签为1和0的占比(也就是概率):

def p(X):
    Fx = X[:,X.shape[1]-1]  # 提取输入样本集的最后一列,也就是所有样本的标签
    x0= 0  # 统计标签为1的样本数量
    x1= 0  # 统计标签为0的样本数量
    for i in range(len(Fx)):
        if Fx[i-1] == 1:  # 这里有个bug!i从0开始时,i-1=-1会取到最后一个元素,应该改成Fx[i]
            x0 = x0+1
        else:
            x1 = x1+1
    P0 = x0/len(Fx)  # 标签为1的概率
    P1 = x1/len(Fx)  # 标签为0的概率
    return(P0,P1)

修正建议:

把循环改成直接遍历标签元素,避免索引错误,代码更直观:

def p(X):
    Fx = X[:,X.shape[1]-1]
    x0, x1 = 0, 0
    for val in Fx:
        if val == 1:
            x0 += 1
        else:
            x1 += 1
    return x0/len(Fx), x1/len(Fx)

用给定的training_set运行的话,标签为1的样本有6个,0的有5个,返回的概率是6/11≈0.545和5/11≈0.455。

3. H(X)函数:计算样本集的信息熵

函数注释里明确提到“needs to be log2”,也就是ID3算法里的信息熵需要用以2为底的对数,但原代码用的是np.log(自然对数),这是需要修正的点,同时还要处理概率为0的边界情况:

def H(X):
    p0, p1 = p(X)
    # 当概率为0时,log2(0)无意义,此时该项对熵的贡献为0,避免报错
    term0 = -p0 * np.log2(p0) if p0 != 0 else 0
    term1 = -p1 * np.log2(p1) if p1 != 0 else 0
    return term0 + term1

熵的公式解释:

信息熵的公式是H = -Σ(p_i * log2(p_i)),其中p_i是每个标签的概率,熵越大说明样本标签越混杂。用刚才的概率计算的话,熵≈-0.545*log2(0.545) -0.455*log2(0.455) ≈0.994,接近1,说明这个样本集确实比较混杂。


内容的提问来源于stack exchange,提问作者Kamil Septio Trojnar

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.22 09:17:49