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.Model;
027import org.tribuo.Trainer;
028import org.tribuo.classification.Label;
029import org.tribuo.classification.LabelFactory;
030import org.tribuo.classification.evaluation.ConfusionMatrix;
031import org.tribuo.classification.evaluation.LabelEvaluation;
032import org.tribuo.classification.evaluation.LabelEvaluator;
033import org.tribuo.data.DataOptions;
034
035import java.io.File;
036import java.io.FileOutputStream;
037import java.io.IOException;
038import java.io.ObjectOutputStream;
039import java.io.OutputStreamWriter;
040import java.io.PrintWriter;
041import java.nio.charset.StandardCharsets;
042import java.nio.file.Paths;
043import java.util.HashMap;
044import java.util.List;
045import java.util.Map;
046import java.util.logging.Level;
047import java.util.logging.Logger;
048
049/**
050 * Trains and tests a model using the supplied data, for each trainer inside a configuration file.
051 */
052public class RunAll {
053    private static final Logger logger = Logger.getLogger(RunAll.class.getName());
054
055    /**
056     * Command line options.
057     */
058    public static class RunAllOptions implements Options {
059        @Override
060        public String getOptionsDescription() {
061            return "Performs the same training and test experiment on all Trainers in the supplied configuration file.";
062        }
063
064        /**
065         * Options for loading in data.
066         */
067        public DataOptions general;
068
069        /**
070         * Directory to write out the models and test reports.
071         */
072        @Option(charName = 'd', longName = "output-directory", usage = "Directory to write out the models and test reports.")
073        public File directory;
074
075        /**
076         * Write out models in protobuf format.
077         */
078        @Option(longName = "write-protobuf-models", usage = "Write out models in protobuf format.")
079        public boolean protobuf;
080    }
081
082    /**
083     * Runs the RunALL CLI.
084     * @param args The CLI arguments.
085     * @throws IOException If it failed to load the data.
086     */
087    public static void main(String[] args) throws IOException {
088        LabsLogFormatter.setAllLogFormatters();
089
090        RunAllOptions o = new RunAllOptions();
091        ConfigurationManager cm;
092        try {
093            cm = new ConfigurationManager(args,o);
094        } catch (UsageException e) {
095            logger.info(e.getMessage());
096            return;
097        }
098
099        if (o.general.trainingPath == null || o.general.testingPath == null || o.directory == null) {
100            logger.info(cm.usage());
101            System.exit(1);
102        }
103        Pair<Dataset<Label>,Dataset<Label>> data = null;
104        try {
105            data = o.general.load(new LabelFactory());
106        } catch (IOException e) {
107            logger.log(Level.SEVERE, "Failed to load data", e);
108            System.exit(1);
109        }
110        Dataset<Label> train = data.getA();
111        Dataset<Label> test = data.getB();
112
113        logger.info("Creating directory - " + o.directory.toString());
114        if (!o.directory.exists() && !o.directory.mkdirs()) {
115            logger.warning("Failed to create directory.");
116        }
117
118        Map<String,Double> performances = new HashMap<>();
119        List<Trainer> trainers = cm.lookupAll(Trainer.class);
120        for (Trainer<?> t : trainers) {
121            String name = t.getClass().getSimpleName();
122            logger.info("Training model using " + t.toString());
123            @SuppressWarnings("unchecked") // configuration system cast.
124            Model<Label> curModel = ((Trainer<Label>)t).train(train);
125            LabelEvaluator evaluator = new LabelEvaluator();
126            LabelEvaluation evaluation = evaluator.evaluate(curModel,test);
127            Double old = performances.put(name,evaluation.microAveragedF1());
128            if (old != null) {
129                logger.info("Found two trainers with the name " + name);
130            }
131            String outputPath = o.directory.toString()+"/"+name;
132            if (o.protobuf) {
133                curModel.serializeToFile(Paths.get(outputPath + ".model"));
134            } else {
135                try (ObjectOutputStream oos = new ObjectOutputStream(new FileOutputStream(outputPath + ".model"))) {
136                    oos.writeObject(curModel);
137                }
138            }
139            try (PrintWriter writer = new PrintWriter(new OutputStreamWriter(new FileOutputStream(outputPath+".output"), StandardCharsets.UTF_8))) {
140                writer.println("Model = " + name);
141                writer.println("Provenance = " + curModel.toString());
142                writer.println();
143                ConfusionMatrix<Label> matrix = evaluation.getConfusionMatrix();
144                writer.println("ConfusionMatrix:\n" + matrix.toString());
145                writer.println();
146                writer.println("Evaluation:\n" + evaluation.toString());
147            }
148        }
149
150        for (Map.Entry<String,Double> e : performances.entrySet()) {
151            logger.info("Trainer = " + e.getKey() + ", F1 = " + e.getValue());
152        }
153
154    }
155
156}