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}