使用Weka+Java结合NBTree与RandomForest时新测试集报错求助
解决NBTree+RandomForest乘积规则组合在新测试集上的ArrayIndexOutOfBoundsException问题
我帮你分析下这个问题:你用乘积规则组合NBTree和RandomForest时,同数据集测试正常,但新数据集测试就抛出数组越界异常,这个问题几乎可以肯定是训练集和测试集的特征结构不匹配导致的,尤其是离散化处理环节出了问题。
先回顾下你的核心代码和错误:
模型组合代码
RandomForest objRandomForest = new RandomForest(); Classifier[] objClassifiers = {objNBTree, objRandomForest}; objVote.setClassifiers(objClassifiers); System.out.println("=== Building Model ==="); objVote.buildClassifier(train); Evaluation eval = new Evaluation(test); eval.evaluateModel(objVote, test); System.out.println(eval.toSummaryString("\nResults\n===============\n", true));
离散化处理代码
BufferedReader reader = new BufferedReader( new FileReader("C:/Users/MUHAMMAD GALA/Documents/CN/Project/NSL_KDD-master/test1000b.arff")); Instances data = new Instances(reader); String[] options = new String[2]; options[0] = "-R"; options[1] = "First-last"; Discretize d = new Discretize(); d.setOptions(options); d.setInputFormat(data); Instances newData = Filter.useFilter(data, d); ArffSaver saver = new ArffSaver(); saver.setInstances(newData); saver.setFile(new File("C:/Users/MUHAMMAD GALA/Documents/CN/Project/NSL_KDD-master/test1000c.arff")); saver.writeBatch();
错误信息
Exception in thread "main" java.lang.ArrayIndexOutOfBoundsException: 12 at weka.classifiers.trees.j48.NBTreeNoSplit.classProb(Unknown Source) at weka.classifiers.trees.j48.ClassifierTree.getProbs(Unknown Source) at weka.classifiers.trees.j48.ClassifierTree.getProbs(Unknown Source) at weka.classifiers.trees.j48.ClassifierTree.distributionForInstance(Unknown Source) at weka.classifiers.trees.NBTree.distributionForInstance(Unknown Source) at weka.classifiers.meta.Vote.distributionForInstanceProduct(Unknown Source) at weka.classifiers.meta.Vote.distributionForInstance(Unknown Source) at weka.classifiers.Evaluation.evaluateModelOnceAndRecordPrediction(Unknown Source) at weka.classifiers.Evaluation.evaluateModel(Unknown Source) at combinationkitesting.CombinationKiTesting.MajorityVotePrediction(CombinationKiTesting.java:88) at combinationkitesting.CombinationKiTesting.main(CombinationKiTesting.java:54)
问题根源
NBTree在训练时会基于训练集离散化后的特征区间计算概率分布,如果你对测试集单独做了离散化(而不是用训练集训练好的过滤器),测试集的特征区间数量、划分规则可能和训练集不一致。比如训练集某个离散化后的属性有10个区间,测试集同属性有13个区间,NBTree访问预存的概率数组时就会出现索引越界。
具体解决方案
1. 共享同一个Discretize过滤器处理训练和测试集
这是最关键的一步:必须先在训练集上训练离散化过滤器,再用这个已经训练好的过滤器处理测试集,保证两者的离散化规则100%一致。修改后的代码示例:
// ---------------------- 处理训练集 ---------------------- BufferedReader trainReader = new BufferedReader( new FileReader("你的训练集路径.arff")); Instances trainData = new Instances(trainReader); trainData.setClassIndex(trainData.numAttributes() - 1); // 统一设置类别索引 // 初始化并训练离散化过滤器 String[] options = new String[2]; options[0] = "-R"; options[1] = "First-last"; Discretize discretizer = new Discretize(); discretizer.setOptions(options); discretizer.setInputFormat(trainData); // 基于训练集设置格式 Instances discretizedTrain = Filter.useFilter(trainData, discretizer); // ---------------------- 处理测试集 ---------------------- BufferedReader testReader = new BufferedReader( new FileReader("你的测试集路径.arff")); Instances testData = new Instances(testReader); testData.setClassIndex(testData.numAttributes() - 1); // 和训练集保持一致 // 关键:用训练好的discretizer处理测试集,而不是重新创建过滤器 Instances discretizedTest = Filter.useFilter(testData, discretizer); // 可选:保存离散化后的测试集 ArffSaver saver = new ArffSaver(); saver.setInstances(discretizedTest); saver.setFile(new File("离散化后的测试集路径.arff")); saver.writeBatch();
2. 验证训练集和测试集的特征一致性
运行代码打印离散化后两者的属性信息,确认每个属性的取值数量完全一致:
// 打印训练集离散化属性 System.out.println("训练集离散化属性信息:"); for (Attribute attr : discretizedTrain.attributes()) { System.out.println(attr.name() + " | 取值数量:" + attr.numValues()); } // 打印测试集离散化属性 System.out.println("\n测试集离散化属性信息:"); for (Attribute attr : discretizedTest.attributes()) { System.out.println(attr.name() + " | 取值数量:" + attr.numValues()); }
如果某个属性的取值数量不一样,说明离散化规则没共享,这就是问题所在。
3. 确保模型训练和测试用的是对齐后的数据集
修改你的模型训练代码,用离散化后的训练集训练,离散化后的测试集测试:
RandomForest objRandomForest = new RandomForest(); Classifier[] objClassifiers = {objNBTree, objRandomForest}; objVote.setClassifiers(objClassifiers); System.out.println("=== Building Model ==="); objVote.buildClassifier(discretizedTrain); // 用离散化训练集 Evaluation eval = new Evaluation(discretizedTest); eval.evaluateModel(objVote, discretizedTest); // 用离散化测试集 System.out.println(eval.toSummaryString("\nResults\n===============\n", true));
按照这个流程调整后,应该就能解决数组越界的问题了。
内容的提问来源于stack exchange,提问作者Anonymous
相关产品推荐
相关产品推荐

