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}