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

Matlab自定义Decision Tree分类器输入参数错误排查求助

自定义MATLAB决策树分类器参数错误修复方案

核心错误分析

你遇到的findBestCutPoint输入参数不足问题,根源是MATLAB类实例方法的参数规则未遵守:所有非静态类方法必须以类实例obj作为第一个输入参数。你的findBestCutPoint和Node方法缺少该参数,导致调用时MATLAB自动传入实例,参数数量不匹配。

此外代码还有几处语法和逻辑错误,以下是逐一修复方案:


1. 修正类方法的参数定义

错误点:

findBestCutPoint和Node方法未将obj作为第一个参数,调用时MATLAB自动补充实例参数,导致参数数量不符。

修复:

  • 对findBestCutPoint添加obj作为第一个参数:
function bestCutPoint = findBestCutPoint(obj, node, X, labels)
  • 推荐将Node改为嵌套类(更符合MATLAB类设计规范,见下文),避免方法调用的参数混乱。

2. 修正变量名不一致问题

错误点:

fit方法中,将findBestCutPoint返回的bestCutPoint错误地写为bestSplit,导致未定义变量,同时rightChild的索引参数也存在错误。

修复:

将fit方法中所有bestSplit替换为bestCutPoint,并修正索引:

bestCutPoint = findBestCutPoint(obj, node, obj.X, obj.Y);
leftChild = my_ClassificationTree.Node(node.startIndex, bestCutPoint.CutIndex - 1);
rightChild = my_ClassificationTree.Node(bestCutPoint.CutIndex, node.endIndex);
obj.numSplits = obj.numSplits + 1;
node.CutPoint = bestCutPoint.CutPoint;
node.CutPredictorIndex = bestCutPoint.CutPredictorIndex;
node.Children = [leftChild, rightChild];

3. 将Node改为嵌套类(规范实现)

当前将Node作为普通方法创建节点对象不符合MATLAB类设计逻辑,改为嵌套类更清晰:
在my_ClassificationTree类内部添加嵌套类定义:

classdef my_ClassificationTree < handle
    properties
        % 原有属性不变
    end
    % 新增嵌套Node类
    classdef Node < handle
        properties
            startIndex
            endIndex
            leaf = false
            Children = []
            size
            CutPoint = 0
            CutPredictorIndex = 0
            NodeClass = 0
        end
        methods
            function obj = Node(sIndex,eIndex)
                obj.startIndex = sIndex;
                obj.endIndex = eIndex;
                obj.size = eIndex - sIndex + 1;
            end
        end
    end
    methods
        % 原有方法...
    end
end

4. 补充实现weightedGDI函数

代码中调用了weightedGDI但未定义,这里补充基于基尼不纯度的加权实现:

function gdi = weightedGDI(nodeLabels, allLabels)
    if isempty(nodeLabels)
        gdi = 0;
        return;
    end
    total = length(allLabels);
    nodeSize = length(nodeLabels);
    classes = unique(nodeLabels);
    gini = 0;
    for c = classes
        p = sum(nodeLabels == c) / nodeSize;
        gini = gini + p^2;
    end
    gdi = (nodeSize / total) * (1 - gini);
end

可将该函数作为独立函数,或添加为my_ClassificationTree的类方法。


5. 修正构造函数的根节点逻辑

原构造函数中重复创建root局部变量,且未将训练后的节点关联到类实例,修复后直接使用obj.root调用fit:

function obj = my_ClassificationTree(X, Y, MinParentSize, MaxNumSplits, Verbose)
    obj.X = X;
    obj.Y = Y;
    obj.MinParentSize = MinParentSize;
    obj.MaxNumSplits = MaxNumSplits;
    obj.Verbose = Verbose;
    obj.numSplits = 0;
    obj.root = my_ClassificationTree.Node(1, size(obj.X,1));
    fit(obj, obj.root);
end

修复后的完整类代码示例

classdef my_ClassificationTree < handle
    properties
        X % training examples
        Y % training labels
        MinParentSize % minimum parent node size
        MaxNumSplits % maximum number of splits        
        Verbose % are we printing out debug as we go?
        numSplits
        root
    end
    
    classdef Node < handle
        properties
            startIndex
            endIndex
            leaf = false
            Children = []
            size
            CutPoint = 0
            CutPredictorIndex = 0
            NodeClass = 0
        end
        methods
            function obj = Node(sIndex,eIndex)
                obj.startIndex = sIndex;
                obj.endIndex = eIndex;
                obj.size = eIndex - sIndex + 1;
            end
        end
    end

    methods
        function obj = my_ClassificationTree(X, Y, MinParentSize, MaxNumSplits, Verbose)
            obj.X = X;
            obj.Y = Y;
            obj.MinParentSize = MinParentSize;
            obj.MaxNumSplits = MaxNumSplits;
            obj.Verbose = Verbose;
            obj.numSplits = 0;
            obj.root = my_ClassificationTree.Node(1, size(obj.X,1));
            fit(obj, obj.root);
        end

        function fit(obj,node)            
            if node.size < obj.MinParentSize || obj.numSplits >= obj.MaxNumSplits
                 node.leaf = true;
                 labels = obj.Y(node.startIndex:node.endIndex);
                 node.NodeClass = mode(labels);
                 return;
            end
            bestCutPoint = findBestCutPoint(obj, node, obj.X, obj.Y);
            leftChild = my_ClassificationTree.Node(node.startIndex, bestCutPoint.CutIndex - 1);
            rightChild = my_ClassificationTree.Node(bestCutPoint.CutIndex, node.endIndex);
            obj.numSplits = obj.numSplits + 1;
            node.CutPoint = bestCutPoint.CutPoint;
            node.CutPredictorIndex = bestCutPoint.CutPredictorIndex;
            node.Children = [leftChild, rightChild];
            fit(obj, leftChild);
            fit(obj, rightChild);
        end        

        function bestCutPoint = findBestCutPoint(obj, node, X, labels)
            bestCutPoint.CutPoint = 0;
            bestCutPoint.CutPredictorIndex = 0;
            bestCutPoint.CutIndex = 0;
            bestGDI = Inf;
            for i = 1:size(X, 2)
                values = unique(X(node.startIndex:node.endIndex, i));
                for j = 1:length(values)
                    currentX = X(node.startIndex:node.endIndex, i);
                    currentLabels = labels(node.startIndex:node.endIndex);
                    leftMask = currentX < values(j);
                    rightMask = currentX >= values(j);
                    leftLabels = currentLabels(leftMask);
                    rightLabels = currentLabels(rightMask);
                    leftGDI = weightedGDI(leftLabels, labels);
                    rightGDI = weightedGDI(rightLabels, labels);
                    cutGDI = leftGDI + rightGDI;
                    if cutGDI < bestGDI
                        bestGDI = cutGDI;
                        bestCutPoint.CutPoint = values(j);
                        bestCutPoint.CutPredictorIndex = i;
                        % 修正CutIndex:取当前节点范围内的第一个匹配索引
                        bestCutPoint.CutIndex = node.startIndex + find(currentX == values(j), 1, 'first') - 1;
                    end
                end
            end
        end

        function predictions = predict(obj, test_examples)
            predictions = categorical(zeros(size(test_examples,1),1));
            for i = 1:size(test_examples, 1)
                currentNode = obj.root;
                while ~currentNode.leaf 
                    value = test_examples(i, currentNode.CutPredictorIndex);
                    if value < currentNode.CutPoint
                        currentNode = currentNode.Children(1);
                    else
                        currentNode = currentNode.Children(2);
                    end
                end
                predictions(i) = currentNode.NodeClass;
            end
        end
    end
end

% 补充weightedGDI函数
function gdi = weightedGDI(nodeLabels, allLabels)
    if isempty(nodeLabels)
        gdi = 0;
        return;
    end
    total = length(allLabels);
    nodeSize = length(nodeLabels);
    classes = unique(nodeLabels);
    gini = 0;
    for c = classes
        p = sum(nodeLabels == c) / nodeSize;
        gini = gini + p^2;
    end
    gdi = (nodeSize / total) * (1 - gini);
end

验证调用

使用你原有的调用代码即可正常运行:

mym2_dt = my_fitctree(train_examples, train_labels, 'MinParentSize', 10)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 00:25:28