如何改进基于基尼不纯度计算信息增益的Python代码?
基于基尼不纯度计算信息增益的代码问题排查与修复
核心错误原因
代码无法通过隐藏测试用例的核心问题是拆分后子集基尼不纯度的加权权重计算逻辑完全错误,另外存在边界场景崩溃、输出格式不符合要求的隐患:
- 权重逻辑错误:基尼不纯度的信息增益计算公式为:
信息增益 = 原集合基尼不纯度 - Σ(单个子集样本数 / 原集合总样本数 * 单个子集基尼不纯度)
原代码错误将权重设置为「子集内1的数量 / 原集合内1的总数量」,只有当两个子集的正样本占比和子集样本量占比恰好相等时结果才会正确,其余场景计算结果完全偏离正确值。 - 除零风险:当原集合全为0、或某个拆分的子集为空时,代码中直接做除法会触发
ZeroDivisionError,直接导致运行失败。 - 输出格式问题:直接使用
round()后打印,当结果末尾为0时不会补全到5位小数(比如结果0.5会打印为0.5而非0.50000),部分严格判题的场景会判定格式错误。
错误复现样例
举一个可复现计算偏差的简单场景:
- 原集合
s = [1,1,0,0],总样本量4,正样本占比0.5,原基尼不纯度为0.5 - 拆分后子集
a = [1,1,0](样本量3,正样本2个,基尼不纯度≈0.44444),子集b = [0](样本量1,正样本0个,基尼不纯度为0) - 正确信息增益:
0.5 - (3/4)*0.44444 - (1/4)*0 ≈ 0.16667 - 原代码计算结果:
0.5 - (2/2)*0.44444 - 0 ≈ 0.05556,和正确值偏差明显。
修复后代码
s = [int(x) for x in input().split()] a = [int(x) for x in input().split()] b = [int(x) for x in input().split()] def calc_gini(arr): # 空集直接返回基尼不纯度0,避免除零错误 n = len(arr) if n == 0: return 0.0 pos_count = sum(1 for num in arr if num == 1) p = pos_count / n return 2 * p * (1 - p) # 计算各集合基尼不纯度 gini_s = calc_gini(s) gini_a = calc_gini(a) gini_b = calc_gini(b) total_len = len(s) # 按标准公式用样本量占比计算加权基尼 if total_len == 0: info_gain = 0.0 else: weighted_gini = (len(a)/total_len)*gini_a + (len(b)/total_len)*gini_b info_gain = gini_s - weighted_gini # 固定输出5位小数,自动补零满足格式要求 print("{0:.5f}".format(info_gain))
修复说明
- 重构基尼计算逻辑,统一在一个函数内完成空集判断、正样本统计、基尼值计算,去掉冗余自定义函数的同时解决除零问题
- 完全修正权重计算逻辑,使用子集样本数/原集合总样本数作为加权系数,符合信息增益的标准定义
- 采用字符串格式化输出固定5位小数,避免浮点数输出格式不符合判题要求的问题
- 覆盖空集、全0/全1集合、子集样本量不等、子集为空等所有边界场景,不会出现运行崩溃问题
内容的提问来源于stack exchange,提问作者Jimmy T.
相关产品推荐
相关产品推荐

