diff --git a/MultiLabel/Core/src/test/java/org/tribuo/multilabel/IndependentMultiLabelTest.java b/MultiLabel/Core/src/test/java/org/tribuo/multilabel/IndependentMultiLabelTest.java index b90105698..0a849465c 100644 --- a/MultiLabel/Core/src/test/java/org/tribuo/multilabel/IndependentMultiLabelTest.java +++ b/MultiLabel/Core/src/test/java/org/tribuo/multilabel/IndependentMultiLabelTest.java @@ -16,23 +16,52 @@ package org.tribuo.multilabel; -import com.oracle.labs.mlrg.olcut.util.Pair; +import java.util.Collections; +import java.util.HashMap; +import java.util.HashSet; +import java.util.Iterator; +import java.util.List; +import java.util.Map; +import java.util.Set; +import java.util.function.Function; +import java.util.logging.Level; +import java.util.logging.Logger; +import java.util.stream.Collectors; +import java.util.stream.StreamSupport; + +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; import org.tribuo.Dataset; +import org.tribuo.Example; +import org.tribuo.ImmutableOutputInfo; import org.tribuo.Model; +import org.tribuo.MutableOutputInfo; +import org.tribuo.OutputInfo; import org.tribuo.Prediction; +import org.tribuo.classification.Label; +import org.tribuo.classification.LabelFactory; +import org.tribuo.classification.evaluation.ClassifierEvaluation; +import org.tribuo.classification.evaluation.ConfusionMatrix; +import org.tribuo.classification.evaluation.ConfusionMetrics; +import org.tribuo.classification.evaluation.LabelConfusionMatrix; +import org.tribuo.classification.evaluation.LabelEvaluation; +import org.tribuo.classification.evaluation.LabelEvaluationUtil; +import org.tribuo.classification.evaluation.LabelMetric; +import org.tribuo.classification.evaluation.LabelMetrics; import org.tribuo.classification.sgd.linear.LinearSGDTrainer; import org.tribuo.classification.sgd.linear.LogisticRegressionTrainer; +import org.tribuo.evaluation.metrics.EvaluationMetric.Average; +import org.tribuo.evaluation.metrics.MetricID; +import org.tribuo.evaluation.metrics.MetricTarget; +import org.tribuo.impl.ArrayExample; import org.tribuo.multilabel.baseline.IndependentMultiLabelTrainer; +import org.tribuo.multilabel.evaluation.MultiLabelEvaluator; import org.tribuo.multilabel.example.MultiLabelDataGenerator; -import org.junit.jupiter.api.Assertions; -import org.junit.jupiter.api.BeforeAll; -import org.junit.jupiter.api.Test; +import org.tribuo.provenance.EvaluationProvenance; import org.tribuo.test.Helpers; - -import java.util.List; -import java.util.Map; -import java.util.logging.Level; -import java.util.logging.Logger; +import com.oracle.labs.mlrg.olcut.util.MutableLong; +import com.oracle.labs.mlrg.olcut.util.Pair; import static org.junit.jupiter.api.Assertions.assertEquals; @@ -67,4 +96,563 @@ public void testIndependentBinaryPredictions() { Helpers.testModelSerialization(model,MultiLabel.class); } + // MultiLabelConfusionMatrix toString() is hard to interpret - convert a MultiLabel evaluation + // to a Label evaluation + @Test + public void multiLabelAsLabel() { + Dataset train = MultiLabelDataGenerator.generateTrainData(); + Dataset test = MultiLabelDataGenerator.generateTestData(); + + IndependentMultiLabelTrainer trainer = new IndependentMultiLabelTrainer( + new LogisticRegressionTrainer()); + Model model = trainer.train(train); + + ClassifierEvaluation evaluation = new MultiLabelEvaluator() + .evaluate(model, test); + + System.out.println(evaluation); + // MultiLabelConfusionMatrix toString() hard to interpret + System.out.println(evaluation.getConfusionMatrix()); + + // given MultiLabelModel model + final List> predictions = asLabelPredictions(evaluation.getPredictions()); + + final LabelConfusionMatrix labelConfusionMatrix = labelConfusionMatrix(model, test); + + // model + final Set labelMetrics = createMetrics(model); + + final Map, Double> results = computeMetrics( + labelConfusionMatrix, labelMetrics, predictions); + + final LabelEvaluation labelEvaluation = labelEvaluation(predictions, results, + labelConfusionMatrix, + model.generatesProbabilities()); + System.out.println(labelEvaluation); + System.out.println(labelEvaluation.getConfusionMatrix()); + } + + private Map, Double> computeMetrics( + final LabelConfusionMatrix labelConfusionMatrix, final Set labelMetrics, + final List> predictions) { + return Collections.unmodifiableMap( + labelMetrics.stream().collect( + Collectors.toMap( + labelMetric -> new MetricID<>(labelMetric.getTarget(), labelMetric.getName()), + labelMetric -> { + final LabelMetrics aLabelMetrics = LabelMetrics + .valueOf(labelMetric.getName()); + final MetricTarget