如何用训练好的Random Forest分类无标签ARFF文件并解决异常?
问题描述
通过提取6波段图像的感兴趣区域及对应坐标标签,创建了带class属性的agricultural.data ARFF数据集:
@RELATION agricultural.data @attribute band0 numeric @attribute band1 numeric @attribute band3 numeric @attribute band4 numeric @attribute band5 numeric @attribute band6 numeric @attribute class {1,2,3,4,5,6,7,8,9} @data -10.95659,-7.61896,-9.8674499,-9.118701,-8.620638,-12.699167,5 ... -9.172866,-9.814803,-10.693634,-13.313326,-8.568673,-12.355089,3
使用该数据集训练Random Forest模型后结果符合预期。
另有一个无class属性的agricultural.data.fullimage ARFF数据集(由完整图像逐行生成):
@RELATION agricultural.data.fullimage @attribute band0 numeric @attribute band1 numeric @attribute band3 numeric @attribute band4 numeric @attribute band5 numeric @attribute band6 numeric @data -9.261405,-7.302625,-10.753542,-8.018068,-7.776727,-12.878252 ... -9.188496,-10.676176,-14.194083,-9.687324,-9.785445,-12.490084
尝试对该数据集分类以实现图像分割,编写代码如下:
FilteredClassifier fc = new FilteredClassifier(); fc.setClassifier(myRandomForestTrainedModel); for(int pixel=0;pixel < ncols;pixel++) { double prediction; /**Some edge case handling**/ prediction = fc.classifyInstance(data.instance(pixel)); //Each data here is a row in the image which I create an ARFF file for byteLinePrediction[pixel] = (byte)Math.floor(prediction+0.5); }
调用classifyInstance时抛出异常:
weka.core.UnassignedClassException: weka.classifiers.meta.FilteredClassifier: Class attribute not set!
无需评估模型性能,仅需生成分割图像,需解决该异常。
解决方法
方案1:直接使用训练好的Random Forest模型(推荐)
如果训练阶段未使用FilteredClassifier,无需额外包裹该类,直接用已训练好的myRandomForestTrainedModel进行分类即可,代码修改如下:
for(int pixel=0;pixel < ncols;pixel++) { double prediction; /**Some edge case handling**/ // 直接调用训练好的Random Forest模型的classifyInstance方法 prediction = myRandomForestTrainedModel.classifyInstance(data.instance(pixel)); byteLinePrediction[pixel] = (byte)Math.floor(prediction+0.5); }
训练好的Random Forest模型已保存了特征与类别映射关系,只要测试数据集的特征顺序、类型与训练集完全一致(本例中6个波段特征完全匹配),即可直接完成分类,无需额外设置类别属性。
方案2:保留FilteredClassifier并设置类别属性
若必须使用FilteredClassifier,需给测试数据集添加与训练集一致的class属性,并设置类别索引,代码示例如下:
import weka.core.Attribute; import java.util.Arrays; // 假设data是加载好的测试数据集 // 添加与训练集一致的class属性 Attribute classAttribute = new Attribute("class", Arrays.asList("1","2","3","4","5","6","7","8","9")); // 将class属性插入到数据集最后一列 data.insertAttributeAt(classAttribute, data.numAttributes()); // 设置数据集的类别索引为最后一列 data.setClassIndex(data.numAttributes() - 1); // 初始化FilteredClassifier并设置模型 FilteredClassifier fc = new FilteredClassifier(); fc.setClassifier(myRandomForestTrainedModel); // 若训练时使用了Filter,需在此同步设置相同的Filter,否则可省略此步 // fc.setFilter(yourTrainingFilter); // 执行分类 for(int pixel=0;pixel < ncols;pixel++) { double prediction; /**Some edge case handling**/ prediction = fc.classifyInstance(data.instance(pixel)); byteLinePrediction[pixel] = (byte)Math.floor(prediction+0.5); }
FilteredClassifier要求明确知晓数据集的类别属性位置,通过添加class属性并设置索引,可满足其运行要求。
内容的提问来源于stack exchange,提问作者Tarun Maganti
相关产品推荐
相关产品推荐

