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}