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.libsvm; 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 libsvm.svm; 024import libsvm.svm_model; 025import libsvm.svm_node; 026import org.tribuo.Example; 027import org.tribuo.ImmutableFeatureMap; 028import org.tribuo.ImmutableOutputInfo; 029import org.tribuo.ONNXExportable; 030import org.tribuo.Prediction; 031import org.tribuo.classification.Label; 032import org.tribuo.classification.libsvm.protos.LibSVMClassificationModelProto; 033import org.tribuo.common.libsvm.KernelType; 034import org.tribuo.common.libsvm.LibSVMModel; 035import org.tribuo.common.libsvm.LibSVMTrainer; 036import org.tribuo.impl.ModelDataCarrier; 037import org.tribuo.protos.core.ModelProto; 038import org.tribuo.provenance.ModelProvenance; 039import org.tribuo.util.Util; 040import org.tribuo.util.onnx.ONNXContext; 041import org.tribuo.util.onnx.ONNXInitializer; 042import org.tribuo.util.onnx.ONNXNode; 043import org.tribuo.util.onnx.ONNXOperators; 044import org.tribuo.util.onnx.ONNXPlaceholder; 045import org.tribuo.util.onnx.ONNXRef; 046 047import java.util.ArrayList; 048import java.util.Arrays; 049import java.util.Collections; 050import java.util.HashMap; 051import java.util.HashSet; 052import java.util.LinkedHashMap; 053import java.util.List; 054import java.util.Map; 055import java.util.Set; 056import java.util.TreeMap; 057import java.util.stream.Collectors; 058 059/** 060 * A classification model that uses an underlying LibSVM model to make the 061 * predictions. 062 * <p> 063 * See: 064 * <pre> 065 * Chang CC, Lin CJ. 066 * "LIBSVM: a library for Support Vector Machines" 067 * ACM transactions on intelligent systems and technology (TIST), 2011. 068 * </pre> 069 * for the nu-svc algorithm: 070 * <pre> 071 * Schölkopf B, Smola A, Williamson R, Bartlett P L. 072 * "New support vector algorithms" 073 * Neural Computation, 2000, 1207-1245. 074 * </pre> 075 * and for the original algorithm: 076 * <pre> 077 * Cortes C, Vapnik V. 078 * "Support-Vector Networks" 079 * Machine Learning, 1995. 080 * </pre> 081 */ 082public class LibSVMClassificationModel extends LibSVMModel<Label> implements ONNXExportable { 083 private static final long serialVersionUID = 3L; 084 085 /** 086 * Protobuf serialization version. 087 */ 088 public static final int CURRENT_VERSION = 0; 089 090 /** 091 * This is used when the model hasn't seen as many outputs as the OutputInfo says are there. 092 * It stores the unseen labels to ensure the predict method has the right number of outputs. 093 * If there are no unobserved labels it's set to Collections.emptySet. 094 */ 095 private final Set<Label> unobservedLabels; 096 097 LibSVMClassificationModel(String name, ModelProvenance description, ImmutableFeatureMap featureIDMap, ImmutableOutputInfo<Label> labelIDMap, List<svm_model> models) { 098 super(name, description, featureIDMap, labelIDMap, models.get(0).param.probability == 1, models); 099 // This sets up the unobservedLabels variable. 100 int[] curLabels = models.get(0).label; 101 if (curLabels.length != labelIDMap.size()) { 102 Map<Integer,Label> tmp = new HashMap<>(); 103 for (Pair<Integer,Label> p : labelIDMap) { 104 tmp.put(p.getA(),p.getB()); 105 } 106 for (int i = 0; i < curLabels.length; i++) { 107 tmp.remove(i); 108 } 109 Set<Label> tmpSet = new HashSet<>(tmp.values().size()); 110 for (Label l : tmp.values()) { 111 tmpSet.add(new Label(l.getLabel(),0.0)); 112 } 113 this.unobservedLabels = Collections.unmodifiableSet(tmpSet); 114 } else { 115 this.unobservedLabels = Collections.emptySet(); 116 } 117 } 118 119 /** 120 * Deserialization factory. 121 * @param version The serialized object version. 122 * @param className The class name. 123 * @param message The serialized data. 124 * @throws InvalidProtocolBufferException If the protobuf could not be parsed from the {@code message}. 125 * @return The deserialized object. 126 */ 127 public static LibSVMClassificationModel deserializeFromProto(int version, String className, Any message) throws InvalidProtocolBufferException { 128 if (version < 0 || version > CURRENT_VERSION) { 129 throw new IllegalArgumentException("Unknown version " + version + ", this class supports at most version " + CURRENT_VERSION); 130 } 131 LibSVMClassificationModelProto proto = message.unpack(LibSVMClassificationModelProto.class); 132 133 ModelDataCarrier<?> carrier = ModelDataCarrier.deserialize(proto.getMetadata()); 134 if (!carrier.outputDomain().getOutput(0).getClass().equals(Label.class)) { 135 throw new IllegalStateException("Invalid protobuf, output domain is not a label domain, found " + carrier.outputDomain().getClass()); 136 } 137 @SuppressWarnings("unchecked") // guarded by getClass 138 ImmutableOutputInfo<Label> outputDomain = (ImmutableOutputInfo<Label>) carrier.outputDomain(); 139 140 svm_model model = deserializeModel(proto.getModel()); 141 142 return new LibSVMClassificationModel(carrier.name(),carrier.provenance(),carrier.featureDomain(),outputDomain,Collections.singletonList(model)); 143 } 144 145 /** 146 * Returns the number of support vectors. 147 * @return The number of support vectors. 148 */ 149 public int getNumberOfSupportVectors() { 150 return models.get(0).SV.length; 151 } 152 153 @Override 154 public Prediction<Label> predict(Example<Label> example) { 155 svm_model model = models.get(0); 156 svm_node[] features = LibSVMTrainer.exampleToNodes(example, featureIDMap, null); 157 // Bias feature is always set 158 if (features.length == 0) { 159 throw new IllegalArgumentException("No features found in Example " + example.toString()); 160 } 161 int[] labels = model.label; 162 double[] scores = new double[labels.length]; 163 if (generatesProbabilities) { 164 svm.svm_predict_probability(model, features, scores); 165 } else { 166 //LibSVM returns a one vs one result, and unpacks it into a score vector by voting 167 double[] onevone = new double[labels.length * (labels.length - 1) / 2]; 168 svm.svm_predict_values(model, features, onevone); 169 int counter = 0; 170 for (int i = 0; i < labels.length; i++) { 171 for (int j = i+1; j < labels.length; j++) { 172 if (onevone[counter] > 0) { 173 scores[i]++; 174 } else { 175 scores[j]++; 176 } 177 counter++; 178 } 179 } 180 } 181 double maxScore = Double.NEGATIVE_INFINITY; 182 Label maxLabel = null; 183 Map<String,Label> map = new LinkedHashMap<>(); 184 for (int i = 0; i < scores.length; i++) { 185 String name = outputIDInfo.getOutput(labels[i]).getLabel(); 186 Label label = new Label(name, scores[i]); 187 map.put(name,label); 188 if (label.getScore() > maxScore) { 189 maxScore = label.getScore(); 190 maxLabel = label; 191 } 192 } 193 if (!unobservedLabels.isEmpty()) { 194 for (Label l : unobservedLabels) { 195 map.put(l.getLabel(),l); 196 } 197 } 198 return new Prediction<>(maxLabel, map, features.length, example, generatesProbabilities); 199 } 200 201 @Override 202 protected LibSVMClassificationModel copy(String newName, ModelProvenance newProvenance) { 203 return new LibSVMClassificationModel(newName,newProvenance,featureIDMap,outputIDInfo,Collections.singletonList(LibSVMModel.copyModel(models.get(0)))); 204 } 205 206 @Override 207 public OnnxMl.ModelProto exportONNXModel(String domain, long modelVersion) { 208 ONNXContext onnx = new ONNXContext(); 209 210 ONNXPlaceholder input = onnx.floatInput(featureIDMap.size()); 211 ONNXPlaceholder output = onnx.floatOutput(outputIDInfo.size()); 212 onnx.setName("Classification-LibSVM"); 213 214 writeONNXGraph(input).assignTo(output); 215 return ONNXExportable.buildModel(onnx, domain, modelVersion, this); 216 } 217 218 @Override 219 public ONNXNode writeONNXGraph(ONNXRef<?> input) { 220 ONNXContext onnx = input.onnxContext(); 221 svm_model model = models.get(0); 222 int numOneVOne = model.label.length * (model.label.length - 1) / 2; 223 int numFeatures = featureIDMap.size(); 224 225 // Extract the attributes 226 Map<String,Object> attributes = new HashMap<>(); 227 attributes.put("classlabels_ints",model.label); 228 float[] coefficients = new float[model.l * (model.nr_class - 1)]; 229 for (int i = 0; i < model.nr_class - 1; i++) { 230 for (int j = 0; j < model.l; j++) { 231 coefficients[i*model.l + j] = (float) model.sv_coef[i][j]; 232 } 233 } 234 attributes.put("coefficients",coefficients); 235 attributes.put("kernel_params",new float[]{(float)model.param.gamma,(float)model.param.coef0,model.param.degree}); 236 attributes.put("kernel_type", KernelType.getKernelType(model.param.kernel_type).name()); 237 float[] rho = new float[model.rho.length]; 238 for (int i = 0; i < rho.length; i++) { 239 rho[i] = (float)-model.rho[i]; 240 } 241 attributes.put("rho",rho); 242 // Extract the support vectors 243 float[] supportVectors = new float[model.l*numFeatures]; 244 245 for (int j = 0; j < model.l; j++) { 246 svm_node[] sv = model.SV[j]; 247 for (svm_node svm_node : sv) { 248 int idx = (j * numFeatures) + svm_node.index; 249 supportVectors[idx] = (float) svm_node.value; 250 } 251 } 252 attributes.put("support_vectors", supportVectors); 253 attributes.put("vectors_per_class", Arrays.copyOf(model.nSV,model.label.length)); 254 if (generatesProbabilities) { 255 attributes.put("prob_a",Arrays.copyOf(Util.toFloatArray(model.probA),numOneVOne)); 256 attributes.put("prob_b",Arrays.copyOf(Util.toFloatArray(model.probB),numOneVOne)); 257 } 258 259 // Build SVM node 260 List<ONNXNode> outputs = input.apply(ONNXOperators.SVM_CLASSIFIER, Arrays.asList("pred_label", "svm_output"), attributes); 261 ONNXNode predLabel = outputs.get(0); 262 ONNXNode svmOutput = outputs.get(1); 263 264 ONNXNode ungatheredOutput = svmOutput; 265 // if the model is not probabilistic we need to vote the one v one classifier output 266 if(!generatesProbabilities) { 267 // If the model has two classes then the scores are inverted for some reason 268 // This is based on the ONNX Runtime behaviour, but the ONNX SVMClassifier spec is ill-defined 269 if(model.nr_class == 2) { 270 ONNXInitializer negOne = onnx.constant("minus_one", -1.0f); 271 ungatheredOutput = writeDecisionFunction(svmOutput.apply(ONNXOperators.MUL, negOne)); 272 } else { 273 ungatheredOutput = writeDecisionFunction(svmOutput); 274 } 275 } 276 277 // Undo the libsvm mapping so the indices line up with Tribuo indices 278 int[] backwardsLibSVMMapping = new int[model.label.length]; 279 for (int i = 0; i < model.label.length; i++) { 280 backwardsLibSVMMapping[model.label[i]] = i; 281 } 282 283 ONNXInitializer indices = onnx.array("label_indices", backwardsLibSVMMapping); 284 285 return ungatheredOutput.apply(ONNXOperators.GATHER, indices, Collections.singletonMap("axis", 1)); 286 } 287 288 private ONNXNode writeDecisionFunction(ONNXNode svmOutputName) { 289 final ONNXContext onnx = svmOutputName.onnxContext(); 290 ONNXInitializer one = onnx.constant("one", 1.0f); 291 ONNXInitializer zero = onnx.constant("zero", 0.0f); 292 293 ONNXNode prediction = svmOutputName.apply(ONNXOperators.LESS, zero).cast(float.class); 294 295 svm_model model = models.get(0); 296 297 TreeMap<Integer, List<ONNXNode>> votes = new TreeMap<>(); 298 299 int k = 0; 300 for (int i = 0; i < model.nr_class; i++) { 301 for (int j = i + 1; j < model.nr_class; j++) { 302 ONNXInitializer index = onnx.constant("Vind_" + k, (long) k); 303 304 ONNXNode extractedFeature = prediction.apply(ONNXOperators.ARRAY_FEATURE_EXTRACTOR, index, "Vsvcv_" + k); 305 votes.computeIfAbsent(j, x -> new ArrayList<>()).add(extractedFeature); 306 307 ONNXNode addNeg = extractedFeature.apply(ONNXOperators.NEG, "Vnegv_" + k).apply(ONNXOperators.ADD, one, "Vnegv1_" + k); 308 votes.computeIfAbsent(i, x -> new ArrayList<>()).add(addNeg); 309 310 k += 1; 311 } 312 } 313 314 List<ONNXNode> oneVOneVotes = votes.values().stream() 315 .map(nodes -> onnx.operation(ONNXOperators.SUM, nodes, "svm_votes")) 316 .collect(Collectors.toList()); 317 /* 318 votes.entrySet().stream().sequential() 319 .sorted(Comparator.comparingInt(Map.Entry::getKey)) 320 .map(Map.Entry::getValue) 321 .map(nodes -> onnx.operation(ONNXOperators.SUM, nodes, "svm_votes")) 322 .collect(Collectors.toList()); 323 324 */ 325 326 return onnx.operation(ONNXOperators.CONCAT, oneVOneVotes, "svm_output", Collections.singletonMap("axis", 1)); 327 } 328 329 @Override 330 public ModelProto serialize() { 331 ModelDataCarrier<Label> carrier = createDataCarrier(); 332 333 LibSVMClassificationModelProto.Builder modelBuilder = LibSVMClassificationModelProto.newBuilder(); 334 modelBuilder.setMetadata(carrier.serialize()); 335 modelBuilder.setModel(serializeModel(models.get(0))); 336 337 ModelProto.Builder builder = ModelProto.newBuilder(); 338 builder.setSerializedData(Any.pack(modelBuilder.build())); 339 builder.setClassName(LibSVMClassificationModel.class.getName()); 340 builder.setVersion(CURRENT_VERSION); 341 342 return builder.build(); 343 } 344}