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.Option; 020import org.tribuo.Trainer; 021import org.tribuo.classification.ClassificationOptions; 022import org.tribuo.classification.Label; 023import org.tribuo.classification.dtree.CARTClassificationOptions; 024import org.tribuo.classification.ensemble.ClassificationEnsembleOptions; 025import org.tribuo.classification.liblinear.LibLinearOptions; 026import org.tribuo.classification.libsvm.LibSVMOptions; 027import org.tribuo.classification.mnb.MultinomialNaiveBayesOptions; 028import org.tribuo.classification.sgd.kernel.KernelSVMOptions; 029import org.tribuo.classification.sgd.linear.LinearSGDOptions; 030import org.tribuo.classification.xgboost.XGBoostOptions; 031import org.tribuo.common.nearest.KNNClassifierOptions; 032import org.tribuo.hash.HashingOptions; 033import org.tribuo.hash.HashingOptions.ModelHashingType; 034 035import java.util.logging.Logger; 036 037/** 038 * Aggregates all the classification algorithms. 039 */ 040public class AllTrainerOptions implements ClassificationOptions<Trainer<Label>> { 041 private static final Logger logger = Logger.getLogger(AllTrainerOptions.class.getName()); 042 043 /** 044 * Types of algorithms supported. 045 */ 046 public enum AlgorithmType { 047 /** 048 * Creates a {@link org.tribuo.classification.dtree.CARTClassificationTrainer}. 049 */ 050 CART, 051 /** 052 * Creates a {@link org.tribuo.common.nearest.KNNTrainer}. 053 */ 054 KNN, 055 /** 056 * Creates a {@link org.tribuo.classification.liblinear.LibLinearClassificationTrainer}. 057 */ 058 LIBLINEAR, 059 /** 060 * Creates a {@link org.tribuo.classification.libsvm.LibSVMClassificationTrainer}. 061 */ 062 LIBSVM, 063 /** 064 * Creates a {@link org.tribuo.classification.mnb.MultinomialNaiveBayesTrainer}. 065 */ 066 MNB, 067 /** 068 * Creates a {@link org.tribuo.classification.sgd.kernel.KernelSVMTrainer}. 069 */ 070 SGD_KERNEL, 071 /** 072 * Creates a {@link org.tribuo.classification.sgd.linear.LinearSGDTrainer}. 073 */ 074 SGD_LINEAR, 075 /** 076 * Creates a {@link org.tribuo.classification.xgboost.XGBoostClassificationTrainer}. 077 */ 078 XGBOOST, 079 } 080 081 /** 082 * Type of learner (or base learner). Defaults to SGD_LINEAR. 083 */ 084 @Option(longName = "algorithm", usage = "Type of learner (or base learner). Defaults to SGD_LINEAR.") 085 public AlgorithmType algorithm = AlgorithmType.SGD_LINEAR; 086 087 /** 088 * Options for CART trainers. 089 */ 090 public CARTClassificationOptions cartOptions; 091 /** 092 * Options for K-NN trainers. 093 */ 094 public KNNClassifierOptions knnOptions; 095 /** 096 * Options for LibLinear trainers. 097 */ 098 public LibLinearOptions liblinearOptions; 099 /** 100 * Options for LibSVM trainers. 101 */ 102 public LibSVMOptions libsvmOptions; 103 /** 104 * Options for Multinomial Naive Bayes trainers. 105 */ 106 public MultinomialNaiveBayesOptions mnbOptions; 107 /** 108 * Options for Kernel SVM trainers. 109 */ 110 public KernelSVMOptions kernelSVMOptions; 111 /** 112 * Options for Linear SGD trainers. 113 */ 114 public LinearSGDOptions linearSGDOptions; 115 /** 116 * Options for XGBoost trainers. 117 */ 118 public XGBoostOptions xgBoostOptions; 119 120 /** 121 * Options for classifier ensembles. 122 */ 123 public ClassificationEnsembleOptions ensemble; 124 /** 125 * Options for hashing trainers. 126 */ 127 public HashingOptions hashingOptions; 128 129 @Override 130 public Trainer<Label> getTrainer() { 131 Trainer<Label> trainer; 132 logger.info("Using " + algorithm); 133 switch (algorithm) { 134 case CART: 135 trainer = cartOptions.getTrainer(); 136 break; 137 case KNN: 138 trainer = knnOptions.getTrainer(); 139 break; 140 case LIBLINEAR: 141 trainer = liblinearOptions.getTrainer(); 142 break; 143 case LIBSVM: 144 trainer = libsvmOptions.getTrainer(); 145 break; 146 case MNB: 147 trainer = mnbOptions.getTrainer(); 148 break; 149 case SGD_KERNEL: 150 trainer = kernelSVMOptions.getTrainer(); 151 break; 152 case SGD_LINEAR: 153 trainer = linearSGDOptions.getTrainer(); 154 break; 155 case XGBOOST: 156 trainer = xgBoostOptions.getTrainer(); 157 break; 158 default: 159 throw new IllegalArgumentException("Unknown classifier " + algorithm); 160 } 161 162 if ((ensemble.ensembleSize > 0) && (ensemble.type != null)) { 163 switch (algorithm) { 164 case XGBOOST: 165 throw new IllegalArgumentException( 166 "Not allowed to ensemble XGBoost models. Why ensemble an ensemble?"); 167 default: 168 trainer = ensemble.wrapTrainer(trainer); 169 break; 170 } 171 } 172 173 if (hashingOptions.modelHashingAlgorithm != ModelHashingType.NONE) { 174 trainer = hashingOptions.getHashedTrainer(trainer); 175 } 176 logger.info("Trainer description " + trainer.toString()); 177 return trainer; 178 } 179 180}