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
相关产品推荐
相关产品推荐

