Matlab自定义ClassificationTree的predict方法每次返回相同类别问题
分类树预测结果全为Setosa的问题排查与修复
核心问题1:Predict方法的叶子节点判断逻辑完全错误
你当前通过遍历Leaves集合、对比训练数据的StartIndex和EndIndex来判断节点是否为叶子,但测试数据与训练数据的索引毫无关联,这个判断逻辑从根本上不成立。更关键的是,你在buildTree中已经给叶子节点设置了node.Leaf = true属性,直接用这个属性判断即可。
修复后的Predict方法:
function predictions = predict(obj,test_exmaples) predictions = categorical; for i = 1:size(test_exmaples,1) currentNode = obj.Root; % 直接通过节点的Leaf属性判断是否到达叶子 while ~currentNode.Leaf if test_exmaples(i,currentNode.CutPredictorIndex) < currentNode.CutPoint currentNode = currentNode.Left; else currentNode = currentNode.Right; end end predictions(i) = currentNode.NodeClass; end end
核心问题2:findBestCutPoint存在索引越界与分割点错误
- 分割点计算索引越界:当
j=1时,values(j-1)会访问无效索引(Matlab数组索引从1开始),导致分割点CutPoint为NaN,预测时任何数值与NaN比较都会返回false,最终一直走右分支或直接返回根节点类别。 - CutIndex查找范围错误:你当前在整个训练集查找
CutIndex,但应该限定在当前节点的索引区间内,否则可能得到超出节点范围的索引,导致子节点创建错误。 - 加权 impurity 计算逻辑冗余:使用当前节点的样本量计算权重即可,无需除以整个数据集的样本量,避免不必要的浮点误差。
修复后的findBestCutPoint关键片段:
for i = 1:size(obj.X,2)% 遍历所有特征 values = unique(obj.X(node.StartIndex:node.EndIndex, i)); values = sort(values); % 从第2个值开始循环,避免j-1索引越界 for j = 2:length(values) leftLabels = obj.Y(node.StartIndex:node.EndIndex, 1); rightLabels = obj.Y(node.StartIndex:node.EndIndex, 1); leftMask = obj.X(node.StartIndex:node.EndIndex, i) < values(j); rightMask = obj.X(node.StartIndex:node.EndIndex, i) >= values(j); leftLabels = leftLabels(leftMask); rightLabels = rightLabels(rightMask); leftSize = length(leftLabels); rightSize = length(rightLabels); % 用当前节点样本量计算权重,而非整个数据集 leftProb = leftSize / node.Size; rightProb = rightSize / node.Size; leftGDI = weightedGDI(leftLabels, obj.Y); rightGDI = weightedGDI(rightLabels, obj.Y); cutGDI = leftProb * leftGDI + rightProb * rightGDI; if cutGDI < bestGDI && leftSize >= minParentSize && rightSize >= minParentSize bestGDI = cutGDI; % 取相邻特征值的中点作为分割点 bestCut.CutPoint = (values(j) + values(j-1))/2; bestCut.CutPredictorIndex = i; % 在当前节点范围内查找索引,转换为全局索引 localIdx = find(obj.X(node.StartIndex:node.EndIndex, i) == values(j), 1, 'first'); bestCut.CutIndex = node.StartIndex + localIdx - 1; end end end
额外验证点
- 确认
Node类的Size属性正确计算为EndIndex - StartIndex + 1,否则可能导致节点被误判为叶子节点不分裂。 - 打印树的结构信息(如根节点的
CutPoint、CutPredictorIndex,左右子节点的Leaf属性),确认树确实按照预期完成了分裂,而非根节点直接作为叶子返回。
内容的提问来源于stack exchange,提问作者BoilingT
相关产品推荐
相关产品推荐

