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}