Weka多层感知器模型CSV输出异常及CSVSaver报错求助
问题解决:Java结合Weka导出预测结果到CSV
问题描述
我用Java结合Weka构建多层感知器模型,代码如下:
import java.io.BufferedWriter; import java.io.FileWriter; import java.util.ArrayList; import weka.classifiers.Evaluation; import weka.classifiers.functions.MultilayerPerceptron; import weka.classifiers.evaluation.Prediction; import weka.core.Instances; import weka.core.converters.ConverterUtils.DataSource; public class multiLayerExample { public static void main(String[] args) throws Exception { String filename = "../train.arff"; DataSource source = new DataSource(filename); Instances train = source.getDataSet(); int cid1 = train.numAttributes() - 1; train.setClassIndex(cid1); Instances validation = DataSource.read("../validation.arff"); int cid2 = validation.numAttributes() - 1; validation.setClassIndex(cid2); Instances test = DataSource.read("../test1.arff"); int cid3 = test.numAttributes() - 1; test.setClassIndex(cid3); MultilayerPerceptron mlp = new MultilayerPerceptron(); mlp.buildClassifier(train); Evaluation eval = new Evaluation(train); eval.evaluateModel(mlp, validation); System.out.println(eval.toSummaryString("\nResults_MLP\n\n", false)); // System.out.println(eval.toClassDetailsString()); // System.out.println(eval.toMatrixString()); ArrayList<Prediction> predictions = eval.predictions(); ArrayList<String[]> predList = new ArrayList<String[]>(predictions.size()); for (int i = 0; i < predictions.size(); i++) { String[] s = new String[1]; s[0] = predictions.get(i).toString(); s[0] = s[0].substring(9, 11); predList.add(s); System.out.println(s); } ArrayList<String[]> li = new ArrayList<String[]>(predictions.size()); li.addAll(predList); System.out.println(li.addAll(predList)); // String csv = "../output.csv"; // CSVSaver writer = new CSVSaver(new FileWriter(csv)); // // writer.writeAll(li); // writer.close(); // //Storing again in csv BufferedWriter writer1 = new BufferedWriter( new FileWriter("../output.csv")); System.out.print(li); writer1.write(li.toString()); writer1.newLine(); writer1.flush(); writer1.close(); } }
遇到两个问题:
- 使用
CSVSaver writer = new CSVSaver(new FileWriter(csv));时,报错“CSVSaver cannot be resolved to a type”; - 改用BufferedWriter写入CSV时,文件输出的是不可读的对象地址格式内容。
需要将数值型预测结果正确写入输出文件。
问题1:CSVSaver无法识别的解决
原因
- 未导入Weka的
CSVSaver类; - Weka的
CSVSaver并不支持直接写入ArrayList<String[]>,它主要用于处理Instances对象。
解决步骤
- 添加导入语句:
import weka.core.converters.CSVSaver; import java.io.File;
- 若要通过
CSVSaver导出结果,需先将预测值封装为Instances对象,但这种方式不如直接用BufferedWriter灵活,更推荐采用问题2的优化方案。
问题2:BufferedWriter输出对象地址的解决
原因
li.toString()会调用ArrayList的默认toString方法,而列表中的元素是String[]数组,数组的toString会输出其内存地址(如[Ljava.lang.String;@xxxxxxx),而非数组内的实际内容。此外,你用substring(9,11)截取预测值的方式不稳定,容易因Prediction.toString()格式变化导致错误。
解决步骤
直接遍历Prediction集合获取预测数值,逐行写入CSV文件,避免中间数组层的冗余。
修正后的完整代码
import java.io.BufferedWriter; import java.io.FileWriter; import java.util.ArrayList; import weka.classifiers.Evaluation; import weka.classifiers.functions.MultilayerPerceptron; import weka.classifiers.evaluation.Prediction; import weka.core.Instances; import weka.core.converters.ConverterUtils.DataSource; public class multiLayerExample { public static void main(String[] args) throws Exception { String filename = "../train.arff"; DataSource source = new DataSource(filename); Instances train = source.getDataSet(); int cid1 = train.numAttributes() - 1; train.setClassIndex(cid1); Instances validation = DataSource.read("../validation.arff"); int cid2 = validation.numAttributes() - 1; validation.setClassIndex(cid2); Instances test = DataSource.read("../test1.arff"); int cid3 = test.numAttributes() - 1; test.setClassIndex(cid3); MultilayerPerceptron mlp = new MultilayerPerceptron(); mlp.buildClassifier(train); Evaluation eval = new Evaluation(train); eval.evaluateModel(mlp, validation); System.out.println(eval.toSummaryString("\nResults_MLP\n\n", false)); ArrayList<Prediction> predictions = eval.predictions(); String csvPath = "../output.csv"; // 使用BufferedWriter正确写入预测结果 try (BufferedWriter writer = new BufferedWriter(new FileWriter(csvPath))) { // 写入表头(可选) writer.write("Predicted_Value"); writer.newLine(); for (Prediction pred : predictions) { // 直接获取预测数值,替代不稳定的字符串截取 double predictedValue = pred.predicted(); // 若为分类任务,可获取对应标签:train.classAttribute().value((int) predictedValue) writer.write(String.valueOf(predictedValue)); writer.newLine(); } } catch (Exception e) { e.printStackTrace(); } } }
关键优化点
- 移除了冗余的
ArrayList<String[]>中间层,直接操作Prediction集合; - 使用
pred.predicted()直接获取预测数值,避免字符串截取的潜在错误; - 采用try-with-resources语法自动关闭流,防止资源泄漏;
- 添加可选表头,让CSV文件更易读。
内容的提问来源于stack exchange,提问作者Micmac
相关产品推荐
相关产品推荐

