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}