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.libsvm;
018
019import com.oracle.labs.mlrg.olcut.config.Config;
020import com.oracle.labs.mlrg.olcut.util.Pair;
021import libsvm.svm;
022import libsvm.svm_model;
023import libsvm.svm_node;
024import libsvm.svm_parameter;
025import libsvm.svm_problem;
026import org.tribuo.Dataset;
027import org.tribuo.Example;
028import org.tribuo.ImmutableFeatureMap;
029import org.tribuo.ImmutableOutputInfo;
030import org.tribuo.Trainer;
031import org.tribuo.classification.Label;
032import org.tribuo.classification.WeightedLabels;
033import org.tribuo.common.libsvm.LibSVMModel;
034import org.tribuo.common.libsvm.LibSVMTrainer;
035import org.tribuo.common.libsvm.SVMParameters;
036import org.tribuo.provenance.ModelProvenance;
037
038import java.util.ArrayList;
039import java.util.Collections;
040import java.util.HashMap;
041import java.util.List;
042import java.util.Map;
043import java.util.SplittableRandom;
044import java.util.logging.Logger;
045
046/**
047 * A trainer for classification models that uses LibSVM.
048 * <p>
049 * Note the train method is synchronized on {@code LibSVMTrainer.class} due to a global RNG in LibSVM.
050 * This is insufficient to ensure reproducibility if LibSVM is used directly in the same JVM as Tribuo, but
051 * avoids locking on classes Tribuo does not control.
052 * <p>
053 * See:
054 * <pre>
055 * Chang CC, Lin CJ.
056 * "LIBSVM: a library for Support Vector Machines"
057 * ACM transactions on intelligent systems and technology (TIST), 2011.
058 * </pre>
059 * for the nu-svc algorithm:
060 * <pre>
061 * Schölkopf B, Smola A, Williamson R, Bartlett P L.
062 * "New support vector algorithms"
063 * Neural Computation, 2000, 1207-1245.
064 * </pre>
065 * and for the original algorithm:
066 * <pre>
067 * Cortes C, Vapnik V.
068 * "Support-Vector Networks"
069 * Machine Learning, 1995.
070 * </pre>
071 */
072public class LibSVMClassificationTrainer extends LibSVMTrainer<Label> implements WeightedLabels {
073    private static final Logger logger = Logger.getLogger(LibSVMClassificationTrainer.class.getName());
074
075    @Config(description="Use Label specific weights.")
076    private Map<String,Float> labelWeights = Collections.emptyMap();
077
078    /**
079     * For OLCUT.
080     */
081    protected LibSVMClassificationTrainer() {}
082
083    /**
084     * Constructs a classification LibSVM trainer using the specified parameters
085     * and {@link Trainer#DEFAULT_SEED}.
086     * @param parameters The SVM parameters.
087     */
088    public LibSVMClassificationTrainer(SVMParameters<Label> parameters) {
089        this(parameters, Trainer.DEFAULT_SEED);
090    }
091
092    /**
093     * Constructs a classification LibSVM trainer using the specified parameters and seed.
094     * @param parameters The SVM parameters.
095     * @param seed The RNG seed for LibSVM's internal RNG.
096     */
097    public LibSVMClassificationTrainer(SVMParameters<Label> parameters, long seed) {
098        super(parameters,seed);
099    }
100
101    /**
102     * Used by the OLCUT configuration system, and should not be called by external code.
103     */
104    @Override
105    public void postConfig() {
106        super.postConfig();
107        if (!svmType.isClassification()) {
108            throw new IllegalArgumentException("Supplied regression or anomaly detection parameters to a classification SVM.");
109        }
110    }
111
112    @Override
113    protected LibSVMModel<Label> createModel(ModelProvenance provenance, ImmutableFeatureMap featureIDMap, ImmutableOutputInfo<Label> outputIDInfo, List<svm_model> models) {
114        return new LibSVMClassificationModel("svm-classification-model", provenance, featureIDMap, outputIDInfo, models);
115    }
116
117    @Override
118    protected List<svm_model> trainModels(svm_parameter curParams, int numFeatures, svm_node[][] features, double[][] outputs, SplittableRandom localRNG) {
119        svm_problem problem = new svm_problem();
120        problem.l = outputs[0].length;
121        problem.x = features;
122        problem.y = outputs[0];
123        if (curParams.gamma == 0) {
124            curParams.gamma = 1.0 / numFeatures;
125        }
126        String checkString = svm.svm_check_parameter(problem, curParams);
127        if(checkString != null) {
128            throw new IllegalArgumentException("Error checking SVM parameters: " + checkString);
129        }
130        // This is safe because we synchronize on LibSVMTrainer.class in the train method to
131        // ensure there is no concurrent use of the rng.
132        svm.rand.setSeed(localRNG.nextLong());
133        return Collections.singletonList(svm.svm_train(problem, curParams));
134    }
135
136    @Override
137    protected Pair<svm_node[][], double[][]> extractData(Dataset<Label> data, ImmutableOutputInfo<Label> outputInfo, ImmutableFeatureMap featureMap) {
138        double[][] ys = new double[1][data.size()];
139        svm_node[][] xs = new svm_node[data.size()][];
140        List<svm_node> buffer = new ArrayList<>();
141        int i = 0;
142        for (Example<Label> example : data) {
143            ys[0][i] = outputInfo.getID(example.getOutput());
144            xs[i] = exampleToNodes(example, featureMap, buffer);
145            i++;
146        }
147        return new Pair<>(xs,ys);
148    }
149
150    @Override
151    protected svm_parameter setupParameters(ImmutableOutputInfo<Label> outputIDInfo) {
152        svm_parameter curParams = SVMParameters.copyParameters(parameters);
153        if (!labelWeights.isEmpty()) {
154            double[] weights = new double[outputIDInfo.size()];
155            int[] indices = new int[outputIDInfo.size()];
156            int i = 0;
157            for (Pair<Integer,Label> label : outputIDInfo) {
158                String labelName = label.getB().getLabel();
159                Float weight = labelWeights.get(labelName);
160                indices[i] = label.getA();
161                if (weight != null) {
162                    weights[i] = weight;
163                } else {
164                    weights[i] = 1.0f;
165                }
166                i++;
167            }
168            curParams.nr_weight = weights.length;
169            curParams.weight = weights;
170            curParams.weight_label = indices;
171            //logger.info("Weights = " + Arrays.toString(weights) + ", labels = " + Arrays.toString(indices) + ", outputIDInfo = " + outputIDInfo);
172        }
173        return curParams;
174    }
175
176    @Override
177    public void setLabelWeights(Map<Label,Float> weights) {
178        labelWeights = new HashMap<>();
179        for (Map.Entry<Label,Float> e : weights.entrySet()) {
180            labelWeights.put(e.getKey().getLabel(),e.getValue());
181        }
182    }
183}