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}