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.liblinear;
018
019import com.oracle.labs.mlrg.olcut.config.Config;
020import com.oracle.labs.mlrg.olcut.util.Pair;
021import org.tribuo.Dataset;
022import org.tribuo.Example;
023import org.tribuo.ImmutableFeatureMap;
024import org.tribuo.ImmutableOutputInfo;
025import org.tribuo.Trainer;
026import org.tribuo.classification.Label;
027import org.tribuo.classification.WeightedLabels;
028import org.tribuo.classification.liblinear.LinearClassificationType.LinearType;
029import org.tribuo.common.liblinear.LibLinearModel;
030import org.tribuo.common.liblinear.LibLinearTrainer;
031import org.tribuo.provenance.ModelProvenance;
032import de.bwaldvogel.liblinear.FeatureNode;
033import de.bwaldvogel.liblinear.Linear;
034import de.bwaldvogel.liblinear.Model;
035import de.bwaldvogel.liblinear.Parameter;
036import de.bwaldvogel.liblinear.Problem;
037
038import java.util.ArrayList;
039import java.util.Collections;
040import java.util.HashMap;
041import java.util.List;
042import java.util.Map;
043import java.util.logging.Logger;
044
045/**
046 * A {@link Trainer} which wraps a liblinear-java classifier trainer.
047 * <p>
048 * See:
049 * <pre>
050 * Fan RE, Chang KW, Hsieh CJ, Wang XR, Lin CJ.
051 * "LIBLINEAR: A library for Large Linear Classification"
052 * Journal of Machine Learning Research, 2008.
053 * </pre>
054 * and for the original algorithm:
055 * <pre>
056 * Cortes C, Vapnik V.
057 * "Support-Vector Networks"
058 * Machine Learning, 1995.
059 * </pre>
060 */
061public class LibLinearClassificationTrainer extends LibLinearTrainer<Label> implements WeightedLabels {
062
063    private static final Logger logger = Logger.getLogger(LibLinearClassificationTrainer.class.getName());
064
065    @Config(description="Use Label specific weights.")
066    private Map<String,Float> labelWeights = Collections.emptyMap();
067
068    /**
069     * Creates a trainer using the default values ({@link LinearType#L2R_L2LOSS_SVC_DUAL}, 1, 0.1, {@link Trainer#DEFAULT_SEED}).
070     */
071    public LibLinearClassificationTrainer() {
072        this(new LinearClassificationType(LinearType.L2R_L2LOSS_SVC_DUAL),1,1000,0.1);
073    }
074
075    /**
076     * Creates a trainer for a LibLinearClassificationModel.
077     * <p>
078     * Uses {@link Trainer#DEFAULT_SEED} as the RNG seed. Sets maxIterations to 1000.
079     * @param trainerType Loss function and optimisation method combination.
080     * @param cost Cost penalty for each incorrectly classified training point.
081     * @param terminationCriterion How close does the optimisation function need to be before terminating that subproblem (usually set to 0.1).
082     */
083    public LibLinearClassificationTrainer(LinearClassificationType trainerType, double cost, double terminationCriterion) {
084        this(trainerType,cost,1000,terminationCriterion);
085    }
086
087    /**
088     * Creates a trainer for a LibLinear model
089     * <p>
090     * Uses {@link Trainer#DEFAULT_SEED} as the RNG seed.
091     * @param trainerType Loss function and optimisation method combination.
092     * @param cost Cost penalty for each incorrectly classified training point.
093     * @param maxIterations The maximum number of dataset iterations.
094     * @param terminationCriterion How close does the optimisation function need to be before terminating that subproblem (usually set to 0.1).
095     */
096    public LibLinearClassificationTrainer(LinearClassificationType trainerType, double cost, int maxIterations, double terminationCriterion) {
097        this(trainerType,cost,maxIterations,terminationCriterion,Trainer.DEFAULT_SEED);
098    }
099
100    /**
101     * Creates a trainer for a LibLinear model
102     * @param trainerType Loss function and optimisation method combination.
103     * @param cost Cost penalty for each incorrectly classified training point.
104     * @param maxIterations The maximum number of dataset iterations.
105     * @param terminationCriterion How close does the optimisation function need to be before terminating that subproblem (usually set to 0.1).
106     * @param seed The RNG seed.
107     */
108    public LibLinearClassificationTrainer(LinearClassificationType trainerType, double cost, int maxIterations, double terminationCriterion, long seed) {
109        super(trainerType,cost,maxIterations,terminationCriterion,seed);
110    }
111
112    /**
113     * Used by the OLCUT configuration system, and should not be called by external code.
114     */
115    @Override
116    public void postConfig() {
117        super.postConfig();
118        if (!trainerType.isClassification()) {
119            throw new IllegalArgumentException("Supplied regression or anomaly detection parameters to a classification linear model.");
120        }
121    }
122
123    @Override
124    protected List<Model> trainModels(Parameter curParams, int numFeatures, FeatureNode[][] features, double[][] outputs) {
125        Problem data = new Problem();
126
127        data.l = features.length;
128        data.y = outputs[0];
129        data.x = features;
130        data.n = numFeatures;
131        data.bias = 1.0;
132
133        return Collections.singletonList(Linear.train(data,curParams));
134    }
135
136    @Override
137    protected LibLinearModel<Label> createModel(ModelProvenance provenance, ImmutableFeatureMap featureIDMap, ImmutableOutputInfo<Label> outputIDInfo, List<Model> models) {
138        if (models.size() != 1) {
139            throw new IllegalArgumentException("Classification uses a single model. Found " + models.size() + " models.");
140        }
141        return new LibLinearClassificationModel("liblinear-classification-model",provenance,featureIDMap,outputIDInfo,models);
142    }
143
144    @Override
145    protected Pair<FeatureNode[][], double[][]> extractData(Dataset<Label> data, ImmutableOutputInfo<Label> outputInfo, ImmutableFeatureMap featureMap) {
146        ArrayList<FeatureNode> featureCache = new ArrayList<>();
147        FeatureNode[][] features = new FeatureNode[data.size()][];
148        double[][] outputs = new double[1][data.size()];
149        int i = 0;
150        for (Example<Label> e : data) {
151            outputs[0][i] = outputInfo.getID(e.getOutput());
152            features[i] = exampleToNodes(e,featureMap,featureCache);
153            i++;
154        }
155        return new Pair<>(features,outputs);
156    }
157
158    @Override
159    protected Parameter setupParameters(ImmutableOutputInfo<Label> labelIDMap) {
160        Parameter curParams = libLinearParams.clone();
161        if (!labelWeights.isEmpty()) {
162            double[] weights = new double[labelIDMap.size()];
163            int[] indices = new int[labelIDMap.size()];
164            int i = 0;
165            for (Pair<Integer,Label> label : labelIDMap) {
166                String labelName = label.getB().getLabel();
167                Float weight = labelWeights.get(labelName);
168                indices[i] = label.getA();
169                if (weight != null) {
170                    weights[i] = weight;
171                } else {
172                    weights[i] = 1.0f;
173                }
174                i++;
175            }
176            curParams.setWeights(weights,indices);
177            //logger.info("Weights = " + Arrays.toString(weights) + ", labels = " + Arrays.toString(indices) + ", outputIDInfo = " + outputIDInfo);
178        }
179        return curParams;
180    }
181
182    @Override
183    public void setLabelWeights(Map<Label,Float> weights) {
184        labelWeights = new HashMap<>();
185        for (Map.Entry<Label,Float> e : weights.entrySet()) {
186            labelWeights.put(e.getKey().getLabel(),e.getValue());
187        }
188    }
189
190}