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 &gt;= 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}