001/*
002 * Copyright (c) 2015-2021, 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.liblinear;
018
019import com.oracle.labs.mlrg.olcut.config.Config;
020import com.oracle.labs.mlrg.olcut.provenance.ConfiguredObjectProvenance;
021import com.oracle.labs.mlrg.olcut.provenance.impl.ConfiguredObjectProvenanceImpl;
022import org.tribuo.classification.Label;
023import org.tribuo.common.liblinear.LibLinearType;
024import de.bwaldvogel.liblinear.SolverType;
025
026import java.io.Serializable;
027
028/**
029 * The carrier type for liblinear classification modes.
030 * <p>
031 * Supports: L1R_L2LOSS_SVC, L2R_L2LOSS_SVC, L2R_L2LOSS_SVC_DUAL, L2R_L1LOSS_SVC_DUAL, MCSVM_CS, L1R_LR, L2R_LR, L2R_LR_DUAL.
032 */
033public final class LinearClassificationType implements LibLinearType<Label> {
034    private static final long serialVersionUID = 1L;
035
036    /**
037     * The different model types available for classification.
038     */
039    public enum LinearType implements Serializable {
040        /**
041         * L1-regularized L2-loss support vector classification
042         */
043        L1R_L2LOSS_SVC(SolverType.L1R_L2LOSS_SVC),
044        /**
045         * L2-regularized L2-loss support vector classification (primal)
046         */
047        L2R_L2LOSS_SVC(SolverType.L2R_L2LOSS_SVC),
048        /**
049         * L2-regularized L2-loss support vector classification (dual)
050         */
051        L2R_L2LOSS_SVC_DUAL(SolverType.L2R_L2LOSS_SVC_DUAL),
052        /**
053         * L2-regularized L1-loss support vector classification (dual)
054         */
055        L2R_L1LOSS_SVC_DUAL(SolverType.L2R_L1LOSS_SVC_DUAL),
056        /**
057         * multi-class support vector classification by Crammer and Singer
058         */
059        MCSVM_CS(SolverType.MCSVM_CS),
060        /**
061         * L1-regularized logistic regression
062         */
063        L1R_LR(SolverType.L1R_LR),
064        /**
065         * L2-regularized logistic regression (primal)
066         */
067        L2R_LR(SolverType.L2R_LR),
068        /**
069         * L2-regularized logistic regression (dual)
070         */
071        L2R_LR_DUAL(SolverType.L2R_LR_DUAL);
072
073        private final SolverType type;
074
075        LinearType(SolverType type) {
076            this.type = type;
077        }
078
079        /**
080         * Gets the LibLinear solver type.
081         * @return The solver type.
082         */
083        public SolverType getSolverType() {
084            return type;
085        }
086    }
087
088    @Config(mandatory=true, description = "The type of classification model")
089    private LinearType type;
090
091    /**
092     * For olcut.
093     */
094    private LinearClassificationType() {}
095
096    /**
097     * Constructs a LinearClassificationType using the supplied algorithm.
098     * @param type The liblinear algorithm.
099     */
100    public LinearClassificationType(LinearType type) {
101        this.type = type;
102    }
103
104    @Override
105    public boolean isClassification() {
106        return true;
107    }
108
109    @Override
110    public boolean isRegression() {
111        return false;
112    }
113
114    @Override
115    public boolean isAnomaly() {
116        return false;
117    }
118
119    @Override
120    public SolverType getSolverType() {
121        return type.getSolverType();
122    }
123
124    @Override
125    public ConfiguredObjectProvenance getProvenance() {
126        return new ConfiguredObjectProvenanceImpl(this,"LibLinearType");
127    }
128
129}