修复DataFrame的索引错误:决策树分类器代码调试求助
嘿,看起来你在手动撸决策树分类器的时候碰到了DataFrame索引的坑,我来帮你捋捋问题出在哪,以及怎么修复:
问题根源分析
你当前代码里的核心错误出在DataFrame的遍历方式上:当你用for entry in data遍历Pandas DataFrame时,迭代的是列名而非行数据!这直接导致后续entry[entry_index]的取值完全不符合预期,进而触发索引错误。另外,通过attrs.index(target)转换列索引的方式也很容易因为列顺序变化出问题。
修复后的完整代码(补全你未写完的部分)
我结合Pandas的特性调整了代码逻辑,同时保留了你原有的决策树核心思路:
import pandas as pd from collections import Counter import math def majority(data, attrs, target): # 获取目标列的众数,作为默认返回值 return data[target].mode()[0] def choose(data, attrs, target): # 用信息增益选择最优划分特征(示例实现,可替换为你的逻辑) def entropy(col): counts = Counter(col) probs = [count/len(col) for count in counts.values()] return -sum(p * math.log2(p) for p in probs if p > 0) base_entropy = entropy(data[target]) best_gain = 0 best_attr = None for attr in attrs: if attr == target: continue # 按特征分组计算条件熵 grouped = data.groupby(attr)[target] conditional_entropy = sum((len(sub)/len(data)) * entropy(sub) for _, sub in grouped) info_gain = base_entropy - conditional_entropy if info_gain > best_gain: best_gain = info_gain best_attr = attr return best_attr def get_vals(data, attrs, pick): # 获取特征的所有唯一取值 return data[pick].unique() def get_data(data, attrs, pick, val): # 根据特征取值筛选子数据集,同时移除当前特征 return data[data[pick] == val].drop(columns=[pick]) def dtree(data, attrs, target): # 空数据集返回众数 if data.empty: return majority(data, attrs, target) # 只剩目标列,返回众数 if len(attrs) - 1 <= 0: return majority(data, attrs, target) # 所有样本目标值一致,直接返回该值 target_vals = data[target].tolist() if len(set(target_vals)) == 1: return target_vals[0] # 选择最优划分特征 pick = choose(data, attrs, target) tree = {pick: {}} # 遍历特征的所有取值,递归构建子树 for val in get_vals(data, attrs, pick): sub_data = get_data(data, attrs, pick, val) sub_attrs = [attr for attr in attrs if attr != pick] tree[pick][val] = dtree(sub_data, sub_attrs, target) return tree
关键修复点说明
- 修正行遍历逻辑:
放弃原有的for entry in data,改用data.itertuples()或data.iterrows()正确获取行数据(示例中直接用Pandas列操作替代了手动遍历取值,更高效)。 - 直接通过列名取值:
移除了attrs.index(target)的索引转换逻辑,直接用列名操作DataFrame,避免了列顺序变化导致的索引混乱。 - 用Pandas内置方法简化数据处理:
用unique()、drop()、groupby()等方法替代手动数据筛选和特征移除,既减少了代码量,也避免了手动操作索引的错误。
你可以把自己实现的majority、choose等函数替换成示例中的对应部分(如果你的逻辑不同),测试后应该就能解决索引错误问题了。
内容的提问来源于stack exchange,提问作者tushariyer
相关产品推荐
相关产品推荐

