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 ai.onnx.proto.OnnxMl; 020import com.google.protobuf.Any; 021import com.google.protobuf.InvalidProtocolBufferException; 022import com.oracle.labs.mlrg.olcut.util.Pair; 023import de.bwaldvogel.liblinear.FeatureNode; 024import de.bwaldvogel.liblinear.Linear; 025import org.tribuo.Example; 026import org.tribuo.Excuse; 027import org.tribuo.Feature; 028import org.tribuo.ImmutableFeatureMap; 029import org.tribuo.ImmutableOutputInfo; 030import org.tribuo.Model; 031import org.tribuo.ONNXExportable; 032import org.tribuo.Prediction; 033import org.tribuo.classification.Label; 034import org.tribuo.common.liblinear.LibLinearModel; 035import org.tribuo.common.liblinear.LibLinearTrainer; 036import org.tribuo.common.liblinear.protos.LibLinearModelProto; 037import org.tribuo.impl.ModelDataCarrier; 038import org.tribuo.provenance.ModelProvenance; 039import org.tribuo.util.onnx.ONNXContext; 040import org.tribuo.util.onnx.ONNXInitializer; 041import org.tribuo.util.onnx.ONNXNode; 042import org.tribuo.util.onnx.ONNXOperators; 043import org.tribuo.util.onnx.ONNXPlaceholder; 044import org.tribuo.util.onnx.ONNXRef; 045 046import java.io.ByteArrayInputStream; 047import java.io.IOException; 048import java.io.ObjectInputStream; 049import java.util.ArrayList; 050import java.util.Arrays; 051import java.util.Collections; 052import java.util.Comparator; 053import java.util.HashMap; 054import java.util.HashSet; 055import java.util.LinkedHashMap; 056import java.util.List; 057import java.util.Map; 058import java.util.PriorityQueue; 059import java.util.Set; 060import java.util.logging.Logger; 061 062/** 063 * A {@link Model} which wraps a LibLinear-java classification model. 064 * <p> 065 * It disables the LibLinear debug output as it's very chatty. 066 * <p> 067 * See: 068 * <pre> 069 * Fan RE, Chang KW, Hsieh CJ, Wang XR, Lin CJ. 070 * "LIBLINEAR: A library for Large Linear Classification" 071 * Journal of Machine Learning Research, 2008. 072 * </pre> 073 * and for the original algorithm: 074 * <pre> 075 * Cortes C, Vapnik V. 076 * "Support-Vector Networks" 077 * Machine Learning, 1995. 078 * </pre> 079 */ 080public class LibLinearClassificationModel extends LibLinearModel<Label> implements ONNXExportable { 081 private static final long serialVersionUID = 3L; 082 083 private static final Logger logger = Logger.getLogger(LibLinearClassificationModel.class.getName()); 084 085 /** 086 * This is used when the model hasn't seen as many outputs as the OutputInfo says are there. 087 * It stores the unseen labels to ensure the predict method has the right number of outputs. 088 * If there are no unobserved labels it's set to Collections.emptySet. 089 */ 090 private final Set<Label> unobservedLabels; 091 092 LibLinearClassificationModel(String name, ModelProvenance description, ImmutableFeatureMap featureIDMap, ImmutableOutputInfo<Label> labelIDMap, List<de.bwaldvogel.liblinear.Model> models) { 093 super(name, description, featureIDMap, labelIDMap, models.get(0).isProbabilityModel(), models); 094 // This sets up the unobservedLabels variable. 095 int[] curLabels = models.get(0).getLabels(); 096 if (curLabels.length != labelIDMap.size()) { 097 Map<Integer,Label> tmp = new HashMap<>(); 098 for (Pair<Integer,Label> p : labelIDMap) { 099 tmp.put(p.getA(),p.getB()); 100 } 101 for (int i = 0; i < curLabels.length; i++) { 102 tmp.remove(i); 103 } 104 Set<Label> tmpSet = new HashSet<>(tmp.values().size()); 105 for (Label l : tmp.values()) { 106 tmpSet.add(new Label(l.getLabel(),0.0)); 107 } 108 this.unobservedLabels = Collections.unmodifiableSet(tmpSet); 109 } else { 110 this.unobservedLabels = Collections.emptySet(); 111 } 112 } 113 114 /** 115 * Deserialization factory. 116 * @param version The serialized object version. 117 * @param className The class name. 118 * @param message The serialized data. 119 * @throws InvalidProtocolBufferException If the protobuf could not be parsed from the {@code message}. 120 * @return The deserialized object. 121 */ 122 public static LibLinearClassificationModel deserializeFromProto(int version, String className, Any message) throws InvalidProtocolBufferException { 123 if (version < 0 || version > CURRENT_VERSION) { 124 throw new IllegalArgumentException("Unknown version " + version + ", this class supports at most version " + CURRENT_VERSION); 125 } 126 if (!"org.tribuo.classification.liblinear.LibLinearClassificationModel".equals(className)) { 127 throw new IllegalStateException("Invalid protobuf, this class can only deserialize LibLinearClassificationModel"); 128 } 129 LibLinearModelProto proto = message.unpack(LibLinearModelProto.class); 130 131 ModelDataCarrier<?> carrier = ModelDataCarrier.deserialize(proto.getMetadata()); 132 if (!carrier.outputDomain().getOutput(0).getClass().equals(Label.class)) { 133 throw new IllegalStateException("Invalid protobuf, output domain is not a label domain, found " + carrier.outputDomain().getClass()); 134 } 135 @SuppressWarnings("unchecked") // guarded by getClass 136 ImmutableOutputInfo<Label> outputDomain = (ImmutableOutputInfo<Label>) carrier.outputDomain(); 137 138 if (proto.getModelsCount() != 1) { 139 throw new IllegalStateException("Invalid protobuf, expected 1 model, found " + proto.getModelsCount()); 140 } 141 try { 142 ByteArrayInputStream bais = new ByteArrayInputStream(proto.getModels(0).toByteArray()); 143 ObjectInputStream ois = new ObjectInputStream(bais); 144 de.bwaldvogel.liblinear.Model model = (de.bwaldvogel.liblinear.Model) ois.readObject(); 145 ois.close(); 146 return new LibLinearClassificationModel(carrier.name(),carrier.provenance(),carrier.featureDomain(),outputDomain,Collections.singletonList(model)); 147 } catch (IOException | ClassNotFoundException e) { 148 throw new IllegalStateException("Invalid protobuf, failed to deserialize liblinear model", e); 149 } 150 } 151 152 @Override 153 public Prediction<Label> predict(Example<Label> example) { 154 FeatureNode[] features = LibLinearTrainer.exampleToNodes(example, featureIDMap, null); 155 // Bias feature is always set 156 if (features.length == 1) { 157 throw new IllegalArgumentException("No features found in Example " + example.toString()); 158 } 159 160 de.bwaldvogel.liblinear.Model model = models.get(0); 161 162 int[] labels = model.getLabels(); 163 double[] scores = new double[labels.length]; 164 165 if (model.isProbabilityModel()) { 166 Linear.predictProbability(model, features, scores); 167 } else { 168 Linear.predictValues(model, features, scores); 169 if ((model.getNrClass() == 2) && (scores[1] == 0.0)) { 170 scores[1] = -scores[0]; 171 } 172 } 173 174 double maxScore = Double.NEGATIVE_INFINITY; 175 Label maxLabel = null; 176 Map<String,Label> map = new LinkedHashMap<>(); 177 for (int i = 0; i < scores.length; i++) { 178 String name = outputIDInfo.getOutput(labels[i]).getLabel(); 179 Label label = new Label(name, scores[i]); 180 map.put(name,label); 181 if (label.getScore() > maxScore) { 182 maxScore = label.getScore(); 183 maxLabel = label; 184 } 185 } 186 if (!unobservedLabels.isEmpty()) { 187 for (Label l : unobservedLabels) { 188 map.put(l.getLabel(),l); 189 } 190 } 191 return new Prediction<>(maxLabel, map, features.length-1, example, generatesProbabilities); 192 } 193 194 @Override 195 public Map<String, List<Pair<String, Double>>> getTopFeatures(int n) { 196 int maxFeatures = n < 0 ? featureIDMap.size() : n; 197 de.bwaldvogel.liblinear.Model model = models.get(0); 198 int[] labels = model.getLabels(); 199 double[] featureWeights = model.getFeatureWeights(); 200 201 Comparator<Pair<String, Double>> comparator = Comparator.comparingDouble(p -> Math.abs(p.getB())); 202 203 /* 204 * Liblinear stores its weights as follows 205 * +------------------+------------------+------------+ 206 * | nr_class weights | nr_class weights | ... 207 * | for 1st feature | for 2nd feature | 208 * +------------------+------------------+------------+ 209 * 210 * If bias >= 0, x becomes [x; bias]. The number of features is 211 * increased by one, so w is a (nr_feature+1)*nr_class array. The 212 * value of bias is stored in the variable bias. 213 */ 214 215 Map<String, List<Pair<String, Double>>> map = new HashMap<>(); 216 int numClasses = model.getNrClass(); 217 int numFeatures = model.getNrFeature(); 218 if (numClasses == 2) { 219 // 220 // When numClasses == 2, liblinear only stores one set of weights. 221 PriorityQueue<Pair<String, Double>> q = new PriorityQueue<>(maxFeatures, comparator); 222 223 for (int i = 0; i < numFeatures; i++) { 224 Pair<String, Double> cur = new Pair<>(featureIDMap.get(i).getName(), featureWeights[i]); 225 if (q.size() < maxFeatures) { 226 q.offer(cur); 227 } else if (comparator.compare(cur, q.peek()) > 0) { 228 q.poll(); 229 q.offer(cur); 230 } 231 } 232 List<Pair<String, Double>> list = new ArrayList<>(); 233 while (q.size() > 0) { 234 list.add(q.poll()); 235 } 236 Collections.reverse(list); 237 map.put(outputIDInfo.getOutput(labels[0]).getLabel(), list); 238 239 List<Pair<String, Double>> otherList = new ArrayList<>(); 240 for (Pair<String, Double> f : list) { 241 Pair<String, Double> otherF = new Pair<>(f.getA(), -f.getB()); 242 otherList.add(otherF); 243 } 244 map.put(outputIDInfo.getOutput(labels[1]).getLabel(), otherList); 245 } else { 246 for (int i = 0; i < labels.length; i++) { 247 PriorityQueue<Pair<String, Double>> q = new PriorityQueue<>(maxFeatures, comparator); 248 //iterate over the non-bias features 249 for (int j = 0; j < numFeatures; j++) { 250 int index = (j * numClasses) + i; 251 Pair<String, Double> cur = new Pair<>(featureIDMap.get(j).getName(), featureWeights[index]); 252 if (q.size() < maxFeatures) { 253 q.offer(cur); 254 } else if (comparator.compare(cur, q.peek()) > 0) { 255 q.poll(); 256 q.offer(cur); 257 } 258 } 259 List<Pair<String, Double>> list = new ArrayList<>(); 260 while (q.size() > 0) { 261 list.add(q.poll()); 262 } 263 Collections.reverse(list); 264 map.put(outputIDInfo.getOutput(labels[i]).getLabel(), list); 265 } 266 } 267 return map; 268 } 269 270 @Override 271 protected LibLinearClassificationModel copy(String newName, ModelProvenance newProvenance) { 272 return new LibLinearClassificationModel(newName,newProvenance,featureIDMap,outputIDInfo,Collections.singletonList(copyModel(models.get(0)))); 273 } 274 275 @Override 276 protected double[][] getFeatureWeights() { 277 double[][] featureWeights = new double[1][]; 278 featureWeights[0] = models.get(0).getFeatureWeights(); 279 return featureWeights; 280 } 281 282 /** 283 * The call to model.getFeatureWeights in the public methods copies the 284 * weights array so this inner method exists to save the copy in getExcuses. 285 * <p> 286 * If it becomes a problem then we could cache the feature weights in the 287 * model. 288 * @param e The example. 289 * @param allFeatureWeights The feature weights. 290 * @return An excuse for this example. 291 */ 292 @Override 293 protected Excuse<Label> innerGetExcuse(Example<Label> e, double[][] allFeatureWeights) { 294 de.bwaldvogel.liblinear.Model model = models.get(0); 295 double[] featureWeights = allFeatureWeights[0]; 296 int[] labels = model.getLabels(); 297 int numClasses = model.getNrClass(); 298 299 Prediction<Label> prediction = predict(e); 300 Map<String, List<Pair<String, Double>>> weightMap = new HashMap<>(); 301 302 if (numClasses == 2) { 303 List<Pair<String, Double>> posScores = new ArrayList<>(); 304 List<Pair<String, Double>> negScores = new ArrayList<>(); 305 for (Feature f : e) { 306 int id = featureIDMap.getID(f.getName()); 307 if (id > -1) { 308 double score = featureWeights[id] * f.getValue(); 309 posScores.add(new Pair<>(f.getName(), score)); 310 negScores.add(new Pair<>(f.getName(), -score)); 311 } 312 } 313 posScores.sort((o1, o2) -> o2.getB().compareTo(o1.getB())); 314 negScores.sort((o1, o2) -> o2.getB().compareTo(o1.getB())); 315 weightMap.put(outputIDInfo.getOutput(labels[0]).getLabel(),posScores); 316 weightMap.put(outputIDInfo.getOutput(labels[1]).getLabel(),negScores); 317 } else { 318 for (int i = 0; i < labels.length; i++) { 319 List<Pair<String, Double>> classScores = new ArrayList<>(); 320 for (Feature f : e) { 321 int id = featureIDMap.getID(f.getName()); 322 if (id > -1) { 323 double score = featureWeights[id * numClasses + i] * f.getValue(); 324 classScores.add(new Pair<>(f.getName(), score)); 325 } 326 } 327 classScores.sort((Pair<String, Double> o1, Pair<String, Double> o2) -> o2.getB().compareTo(o1.getB())); 328 weightMap.put(outputIDInfo.getOutput(labels[i]).getLabel(), classScores); 329 } 330 } 331 332 return new Excuse<>(e, prediction, weightMap); 333 } 334 335 @Override 336 public OnnxMl.ModelProto exportONNXModel(String domain, long modelVersion) { 337 ONNXContext onnx = new ONNXContext(); 338 339 onnx.setName("Classification-LibLinear"); 340 ONNXPlaceholder input = onnx.floatInput(featureIDMap.size()); 341 ONNXPlaceholder output = onnx.floatOutput(outputIDInfo.size()); 342 343 // Build graph 344 writeONNXGraph(input).assignTo(output); 345 346 return ONNXExportable.buildModel(onnx, domain, modelVersion, this); 347 } 348 349 @Override 350 public ONNXNode writeONNXGraph(ONNXRef<?> input) { 351 352 ONNXContext onnx = input.onnxContext(); 353 354 de.bwaldvogel.liblinear.Model model = models.get(0); 355 double[] rawWeights = model.getFeatureWeights(); 356 int[] labels = model.getLabels(); 357 int numFeatures = featureIDMap.size(); 358 int numLabels = labels.length; 359 if (numLabels != outputIDInfo.size()) { 360 throw new IllegalStateException("Unexpected number of labels, output domain = " + outputIDInfo.size() + ", LibLinear's internal count = " + numLabels); 361 } 362 363 // setup weight arrays for easy processing 364 if (model.getNrClass() == 2) { 365 // Replicate weights in binary problems 366 double[] newWeights = new double[rawWeights.length*2]; 367 for (int i = 0; i < rawWeights.length; i++) { 368 if (labels[0] == 0) { 369 newWeights[i * 2] = rawWeights[i]; 370 newWeights[(i * 2) + 1] = -rawWeights[i]; 371 } else { 372 newWeights[i * 2] = -rawWeights[i]; 373 newWeights[(i * 2) + 1] = rawWeights[i]; 374 } 375 } 376 rawWeights = newWeights; 377 } else { 378 double[] newWeights = new double[rawWeights.length]; 379 for (int j = 0; j < numFeatures + 1; j++) { 380 for (int i = 0; i < numLabels; i++) { 381 int newIdx = (j * numLabels) + labels[i]; 382 int oldIdx = (j * numLabels) + i; 383 newWeights[newIdx] = rawWeights[oldIdx]; 384 } 385 } 386 rawWeights = newWeights; 387 } 388 389 final double[] weights = rawWeights; 390 391 ONNXInitializer weightTensor = onnx.floatTensor("liblinear_weights", Arrays.asList(numFeatures, numLabels), fb -> { 392 for (int i = 0; i < weights.length - numLabels; i++) { 393 fb.put((float) weights[i]); 394 } 395 }); 396 397 ONNXInitializer biasTensor = onnx.floatTensor("liblinear_biases", Collections.singletonList(numLabels), fb -> { 398 for (int i = numFeatures * numLabels; i < weights.length; i++) { 399 fb.put((float) weights[i]); 400 } 401 }); 402 403 ONNXNode gemm = input.apply(ONNXOperators.GEMM, Arrays.asList(weightTensor, biasTensor)); 404 405 if(model.isProbabilityModel()) { 406 return gemm.apply(ONNXOperators.SOFTMAX, Collections.singletonMap("axis", 1)); 407 } else { 408 return gemm; 409 } 410 } 411 412}