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

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存在索引越界与分割点错误

  1. 分割点计算索引越界:当j=1时,values(j-1)会访问无效索引(Matlab数组索引从1开始),导致分割点CutPoint为NaN,预测时任何数值与NaN比较都会返回false,最终一直走右分支或直接返回根节点类别。
  2. CutIndex查找范围错误:你当前在整个训练集查找CutIndex,但应该限定在当前节点的索引区间内,否则可能得到超出节点范围的索引,导致子节点创建错误。
  3. 加权 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

额外验证点

  1. 确认Node类的Size属性正确计算为EndIndex - StartIndex + 1,否则可能导致节点被误判为叶子节点不分裂。
  2. 打印树的结构信息(如根节点的CutPoint、CutPredictorIndex,左右子节点的Leaf属性),确认树确实按照预期完成了分裂,而非根节点直接作为叶子返回。

内容的提问来源于stack exchange,提问作者BoilingT

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 14:15:48