001/* 002 * Copyright (c) 2015, 2022, Oracle and/or its affiliates. All rights reserved. 003 * 004 * Licensed under the Apache License, Version 2.0 (the "License"); 005 * you may not use this file except in compliance with the License. 006 * You may obtain a copy of the License at 007 * 008 * http://www.apache.org/licenses/LICENSE-2.0 009 * 010 * Unless required by applicable law or agreed to in writing, software 011 * distributed under the License is distributed on an "AS IS" BASIS, 012 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express implied. 013 * See the License for the specific language governing permissions and 014 * limitations under the License. 015 */ 016 017package org.tribuo.classification.experiments; 018 019import com.oracle.labs.mlrg.olcut.config.ConfigurationManager; 020import com.oracle.labs.mlrg.olcut.config.Option; 021import com.oracle.labs.mlrg.olcut.config.Options; 022import com.oracle.labs.mlrg.olcut.config.UsageException; 023import com.oracle.labs.mlrg.olcut.util.LabsLogFormatter; 024import com.oracle.labs.mlrg.olcut.util.Pair; 025import org.tribuo.Dataset; 026import org.tribuo.Example; 027import org.tribuo.ImmutableDataset; 028import org.tribuo.Model; 029import org.tribuo.Prediction; 030import org.tribuo.classification.Label; 031import org.tribuo.classification.LabelFactory; 032import org.tribuo.classification.evaluation.LabelEvaluation; 033import org.tribuo.classification.evaluation.LabelEvaluator; 034import org.tribuo.data.DataOptions; 035import org.tribuo.data.csv.CSVLoader; 036import org.tribuo.data.text.TextDataSource; 037import org.tribuo.data.text.TextFeatureExtractor; 038import org.tribuo.data.text.impl.SimpleTextDataSource; 039import org.tribuo.data.text.impl.TextFeatureExtractorImpl; 040import org.tribuo.data.text.impl.TokenPipeline; 041import org.tribuo.datasource.LibSVMDataSource; 042import org.tribuo.util.Util; 043import org.tribuo.util.tokens.impl.BreakIteratorTokenizer; 044 045import java.io.BufferedInputStream; 046import java.io.BufferedWriter; 047import java.io.FileInputStream; 048import java.io.IOException; 049import java.io.ObjectInputStream; 050import java.nio.file.Files; 051import java.nio.file.Path; 052import java.util.List; 053import java.util.Locale; 054import java.util.logging.Level; 055import java.util.logging.Logger; 056import java.util.stream.Collectors; 057 058/** 059 * Test a classifier for a standard dataset. 060 */ 061public class Test { 062 063 private static final Logger logger = Logger.getLogger(Test.class.getName()); 064 065 /** 066 * Command line options. 067 */ 068 public static class ConfigurableTestOptions implements Options { 069 @Override 070 public String getOptionsDescription() { 071 return "Tests an already trained classifier on a dataset."; 072 } 073 074 /** 075 * Hashing dimension used for standard text format. 076 */ 077 @Option(longName = "hashing-dimension", usage = "Hashing dimension used for standard text format.") 078 public int hashDim = 0; 079 /** 080 * Ngram size to generate when using standard text format. Defaults to 2. 081 */ 082 @Option(longName = "ngram", usage = "Ngram size to generate when using standard text format. Defaults to 2.") 083 public int ngram = 2; 084 /** 085 * Use term counts instead of boolean when using the standard text format. 086 */ 087 @Option(longName = "term-counting", usage = "Use term counts instead of boolean when using the standard text format.") 088 public boolean termCounting; 089 /** 090 * Response name in the csv file. 091 */ 092 @Option(longName = "csv-response-name", usage = "Response name in the csv file.") 093 public String csvResponseName; 094 /** 095 * Is the libsvm file zero indexed. 096 */ 097 @Option(longName = "libsvm-zero-indexed", usage = "Is the libsvm file zero indexed.") 098 public boolean zeroIndexed = false; 099 /** 100 * Load a trainer from the config file. 101 */ 102 @Option(charName = 'f', longName = "model-path", usage = "Load a trainer from the config file.") 103 public Path modelPath; 104 /** 105 * Path to write model predictions 106 */ 107 @Option(charName = 'o', longName = "predictions", usage = "Path to write model predictions") 108 public Path predictionPath; 109 /** 110 * Loads the data using the specified format. Defaults to LIBSVM. 111 */ 112 @Option(charName = 's', longName = "input-format", usage = "Loads the data using the specified format. Defaults to LIBSVM.") 113 public DataOptions.InputFormat inputFormat = DataOptions.InputFormat.LIBSVM; 114 /** 115 * Path to the testing file. 116 */ 117 @Option(charName = 'v', longName = "testing-file", usage = "Path to the testing file.") 118 public Path testingPath; 119 /** 120 * Load the model in protobuf format. 121 */ 122 @Option(longName = "read-protobuf-model", usage = "Load the model in protobuf format.") 123 public boolean protobufModel; 124 } 125 126 /** 127 * Loads in the model and the dataset from the options. 128 * @param o The options. 129 * @return The model and the dataset. 130 * @throws IOException If either the model or dataset could not be read. 131 */ 132 @SuppressWarnings("unchecked") // deserialising generically typed datasets. 133 public static Pair<Model<Label>,Dataset<Label>> load(ConfigurableTestOptions o) throws IOException { 134 Path modelPath = o.modelPath; 135 Path datasetPath = o.testingPath; 136 logger.info(String.format("Loading model from %s", modelPath)); 137 Model<?> tmpModel; 138 if (o.protobufModel) { 139 tmpModel = Model.deserializeFromFile(modelPath); 140 } else { 141 try (ObjectInputStream mois = new ObjectInputStream(new BufferedInputStream(new FileInputStream(modelPath.toFile())))) { 142 tmpModel = (Model<?>) mois.readObject(); 143 } catch (ClassNotFoundException e) { 144 throw new IllegalArgumentException("Unknown class in serialised model", e); 145 } 146 } 147 Model<Label> model = tmpModel.castModel(Label.class); 148 logger.info(String.format("Loading data from %s", datasetPath)); 149 Dataset<Label> test; 150 switch (o.inputFormat) { 151 case SERIALIZED: 152 // 153 // Load Tribuo serialised datasets. 154 logger.info("Deserialising dataset from " + datasetPath); 155 try (ObjectInputStream oits = new ObjectInputStream(new BufferedInputStream(new FileInputStream(datasetPath.toFile())))) { 156 Dataset<Label> deserTest = (Dataset<Label>) oits.readObject(); 157 test = ImmutableDataset.copyDataset(deserTest,model.getFeatureIDMap(),model.getOutputIDInfo()); 158 logger.info(String.format("Loaded %d testing examples for %s", test.size(), test.getOutputs().toString())); 159 } catch (ClassNotFoundException e) { 160 throw new IllegalArgumentException("Unknown class in serialised dataset", e); 161 } 162 break; 163 case SERIALIZED_PROTOBUF: 164 // 165 // Load Tribuo protobuf serialised datasets. 166 Dataset<?> tmp = Dataset.deserializeFromFile(datasetPath); 167 if (tmp.validate(Label.class)) { 168 test = Dataset.castDataset(tmp, Label.class); 169 test = ImmutableDataset.copyDataset(test,model.getFeatureIDMap(),model.getOutputIDInfo()); 170 logger.info(String.format("Loaded %d testing examples for %s", test.size(), test.getOutputs().toString())); 171 } else { 172 throw new IllegalArgumentException("Invalid test dataset type, expected Label.class"); 173 } 174 break; 175 case LIBSVM: 176 // 177 // Load the libsvm text-based data format. 178 boolean zeroIndexed = o.zeroIndexed; 179 int maxFeatureID = model.getFeatureIDMap().size() - 1; 180 LibSVMDataSource<Label> testSVMSource = new LibSVMDataSource<>(datasetPath,new LabelFactory(),zeroIndexed,maxFeatureID); 181 test = new ImmutableDataset<>(testSVMSource,model,true); 182 logger.info(String.format("Loaded %d training examples for %s", test.size(), test.getOutputs().toString())); 183 break; 184 case TEXT: 185 // 186 // Using a simple Java break iterator to generate ngram features. 187 TextFeatureExtractor<Label> extractor; 188 if (o.hashDim > 0) { 189 extractor = new TextFeatureExtractorImpl<>(new TokenPipeline(new BreakIteratorTokenizer(Locale.US), o.ngram, o.termCounting, o.hashDim)); 190 } else { 191 extractor = new TextFeatureExtractorImpl<>(new TokenPipeline(new BreakIteratorTokenizer(Locale.US), o.ngram, o.termCounting)); 192 } 193 194 TextDataSource<Label> testSource = new SimpleTextDataSource<>(datasetPath, new LabelFactory(), extractor); 195 test = new ImmutableDataset<>(testSource, model.getFeatureIDMap(), model.getOutputIDInfo(),true); 196 logger.info(String.format("Loaded %d testing examples for %s", test.size(), test.getOutputs().toString())); 197 break; 198 case CSV: 199 // 200 // Load the data using the simple CSV loader 201 if (o.csvResponseName == null) { 202 throw new IllegalArgumentException("Please supply a response column name"); 203 } 204 CSVLoader<Label> loader = new CSVLoader<>(new LabelFactory()); 205 test = new ImmutableDataset<>(loader.loadDataSource(datasetPath,o.csvResponseName),model.getFeatureIDMap(),model.getOutputIDInfo(),true); 206 logger.info(String.format("Loaded %d testing examples for %s", test.size(), test.getOutputs().toString())); 207 break; 208 default: 209 throw new IllegalArgumentException("Unsupported input format " + o.inputFormat); 210 } 211 return new Pair<>(model,test); 212 } 213 214 /** 215 * Runs the Test CLI. 216 * @param args the command line arguments 217 */ 218 public static void main(String[] args) { 219 220 // 221 // Use the labs format logging. 222 LabsLogFormatter.setAllLogFormatters(); 223 224 ConfigurableTestOptions o = new ConfigurableTestOptions(); 225 ConfigurationManager cm; 226 try { 227 cm = new ConfigurationManager(args,o); 228 } catch (UsageException e) { 229 logger.info(e.getMessage()); 230 return; 231 } 232 233 if (o.modelPath == null || o.testingPath == null) { 234 logger.info(cm.usage()); 235 System.exit(1); 236 } 237 Pair<Model<Label>,Dataset<Label>> loaded = null; 238 try { 239 loaded = load(o); 240 } catch (IOException e) { 241 logger.log(Level.SEVERE, "Failed to load model/data", e); 242 System.exit(1); 243 } 244 Model<Label> model = loaded.getA(); 245 Dataset<Label> test = loaded.getB(); 246 247 logger.info("Model is " + model.toString()); 248 logger.info("Labels are " + model.getOutputIDInfo().toReadableString()); 249 250 LabelEvaluator labelEvaluator = new LabelEvaluator(); 251 final long testStart = System.currentTimeMillis(); 252 List<Prediction<Label>> predictions = model.predict(test); 253 LabelEvaluation evaluation = labelEvaluator.evaluate(model,predictions,test.getProvenance()); 254 final long testStop = System.currentTimeMillis(); 255 logger.info("Finished evaluating model " + Util.formatDuration(testStart,testStop)); 256 System.out.println(evaluation.toString()); 257 System.out.println(evaluation.getConfusionMatrix().toString()); 258 if (model.generatesProbabilities()) { 259 System.out.println("Average AUC = " + evaluation.averageAUCROC(false)); 260 System.out.println("Average weighted AUC = " + evaluation.averageAUCROC(true)); 261 } 262 263 if (o.predictionPath!=null) { 264 try(BufferedWriter wrt = Files.newBufferedWriter(o.predictionPath)) { 265 List<String> labels = model.getOutputIDInfo().getDomain().stream().map(Label::getLabel).sorted().collect(Collectors.toList()); 266 wrt.write("Label,"); 267 wrt.write(String.join(",", labels)); 268 wrt.newLine(); 269 for(Prediction<Label> pred : predictions) { 270 Example<Label> ex = pred.getExample(); 271 wrt.write(ex.getOutput().getLabel()+","); 272 wrt.write(labels 273 .stream() 274 .map(l -> Double.toString(pred 275 .getOutputScores() 276 .get(l).getScore())) 277 .collect(Collectors.joining(","))); 278 wrt.newLine(); 279 } 280 wrt.flush(); 281 } catch (IOException e) { 282 logger.log(Level.SEVERE, "Error writing predictions", e); 283 } 284 } 285 286 } 287 288}