001/* 002 * Copyright (c) 2015-2020, 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.Model; 028import org.tribuo.MutableDataset; 029import org.tribuo.Prediction; 030import org.tribuo.Trainer; 031import org.tribuo.WeightedExamples; 032import org.tribuo.classification.Label; 033import org.tribuo.classification.LabelFactory; 034import org.tribuo.classification.WeightedLabels; 035import org.tribuo.classification.evaluation.ConfusionMatrix; 036import org.tribuo.classification.evaluation.LabelEvaluation; 037import org.tribuo.classification.evaluation.LabelEvaluator; 038import org.tribuo.data.DataOptions; 039import org.tribuo.util.Util; 040 041import java.io.BufferedWriter; 042import java.io.IOException; 043import java.nio.file.Files; 044import java.nio.file.Path; 045import java.util.HashMap; 046import java.util.List; 047import java.util.Map; 048import java.util.logging.Level; 049import java.util.logging.Logger; 050import java.util.stream.Collectors; 051 052/** 053 * Build and run a classifier for a standard dataset. 054 */ 055public class ConfigurableTrainTest { 056 057 private static final Logger logger = Logger.getLogger(ConfigurableTrainTest.class.getName()); 058 059 /** 060 * Command line options. 061 */ 062 public static class ConfigurableTrainTestOptions implements Options { 063 @Override 064 public String getOptionsDescription() { 065 return "Loads a Trainer (and optionally a Datasource) from a config file, trains a Model, tests it and optionally saves it to disk."; 066 } 067 068 /** 069 * Options for loading in data. 070 */ 071 public DataOptions general; 072 073 /** 074 * Load a trainer from the config file. 075 */ 076 @Option(charName = 't', longName = "trainer", usage = "Load a trainer from the config file.") 077 public Trainer<Label> trainer; 078 079 /** 080 * A list of weights to use in classification. Format = LABEL_NAME:weight,LABEL_NAME:weight... 081 */ 082 @Option(charName = 'w', longName = "weights", usage = "A list of weights to use in classification. Format = LABEL_NAME:weight,LABEL_NAME:weight...") 083 public List<String> weights; 084 085 /** 086 * Path to write model predictions 087 */ 088 @Option(charName = 'o', longName = "predictions", usage = "Path to write model predictions") 089 public Path predictionPath; 090 } 091 092 /** 093 * Converts the weight text input format into an object suitable for use in a Trainer. 094 * @param input The input form. 095 * @return The weights. 096 */ 097 public static Map<Label,Float> processWeights(List<String> input) { 098 Map<Label,Float> map = new HashMap<>(); 099 100 for (String tuple : input) { 101 String[] splitTuple = tuple.split(":"); 102 map.put(new Label(splitTuple[0]),Float.parseFloat(splitTuple[1])); 103 } 104 105 return map; 106 } 107 108 /** 109 * @param args the command line arguments 110 */ 111 public static void main(String[] args) { 112 113 // 114 // Use the labs format logging. 115 LabsLogFormatter.setAllLogFormatters(); 116 117 ConfigurableTrainTestOptions o = new ConfigurableTrainTestOptions(); 118 ConfigurationManager cm; 119 try { 120 cm = new ConfigurationManager(args,o); 121 } catch (UsageException e) { 122 logger.info(e.getMessage()); 123 return; 124 } 125 126 if (o.general.trainingPath == null || o.general.testingPath == null) { 127 logger.info(cm.usage()); 128 System.exit(1); 129 } 130 Pair<Dataset<Label>,Dataset<Label>> data = null; 131 try { 132 data = o.general.load(new LabelFactory()); 133 } catch (IOException e) { 134 logger.log(Level.SEVERE, "Failed to load data", e); 135 System.exit(1); 136 } 137 Dataset<Label> train = data.getA(); 138 Dataset<Label> test = data.getB(); 139 140 if (o.trainer == null) { 141 logger.warning("No trainer supplied"); 142 logger.info(cm.usage()); 143 System.exit(1); 144 } 145 logger.info("Trainer is " + o.trainer.toString()); 146 147 if (o.weights != null) { 148 Map<Label,Float> weightsMap = processWeights(o.weights); 149 if (o.trainer instanceof WeightedLabels) { 150 ((WeightedLabels) o.trainer).setLabelWeights(weightsMap); 151 logger.info("Setting label weights using " + weightsMap.toString()); 152 } else if (o.trainer instanceof WeightedExamples) { 153 ((MutableDataset<Label>)train).setWeights(weightsMap); 154 logger.info("Setting example weights using " + weightsMap.toString()); 155 } else { 156 logger.warning("The selected trainer does not support weighted training. The chosen trainer is " + o.trainer.toString()); 157 logger.info(cm.usage()); 158 System.exit(1); 159 } 160 } 161 162 logger.info("Labels are " + train.getOutputInfo().toReadableString()); 163 164 final long trainStart = System.currentTimeMillis(); 165 Model<Label> model = o.trainer.train(train); 166 final long trainStop = System.currentTimeMillis(); 167 168 logger.info("Finished training classifier " + Util.formatDuration(trainStart,trainStop)); 169 170 LabelEvaluator labelEvaluator = new LabelEvaluator(); 171 final long testStart = System.currentTimeMillis(); 172 List<Prediction<Label>> predictions = model.predict(test); 173 LabelEvaluation labelEvaluation = labelEvaluator.evaluate(model,predictions,test.getProvenance()); 174 final long testStop = System.currentTimeMillis(); 175 logger.info("Finished evaluating model " + Util.formatDuration(testStart,testStop)); 176 System.out.println(labelEvaluation.toString()); 177 ConfusionMatrix<Label> matrix = labelEvaluation.getConfusionMatrix(); 178 System.out.println(matrix.toString()); 179 if (model.generatesProbabilities()) { 180 System.out.println("Average AUC = " + labelEvaluation.averageAUCROC(false)); 181 System.out.println("Average weighted AUC = " + labelEvaluation.averageAUCROC(true)); 182 } 183 184 if(o.predictionPath!=null) { 185 try(BufferedWriter wrt = Files.newBufferedWriter(o.predictionPath)) { 186 List<String> labels = model.getOutputIDInfo().getDomain().stream().map(Label::getLabel).sorted().collect(Collectors.toList()); 187 wrt.write("Label,"); 188 wrt.write(String.join(",", labels)); 189 wrt.newLine(); 190 for(Prediction<Label> pred : predictions) { 191 Example<Label> ex = pred.getExample(); 192 wrt.write(ex.getOutput().getLabel()+","); 193 wrt.write(labels 194 .stream() 195 .map(l -> Double.toString(pred 196 .getOutputScores() 197 .get(l).getScore())) 198 .collect(Collectors.joining(","))); 199 wrt.newLine(); 200 } 201 wrt.flush(); 202 } catch (IOException e) { 203 logger.log(Level.SEVERE, "Error writing predictions", e); 204 } 205 } 206 207 if (o.general.outputPath != null) { 208 try { 209 o.general.saveModel(model); 210 } catch (IOException e) { 211 logger.log(Level.SEVERE, "Error writing model", e); 212 } 213 } 214 } 215}