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}