66using Microsoft . ML . Runtime . Api ;
77using Microsoft . ML . Runtime . Data ;
88using System ;
9+ using System . Collections . Generic ;
910
1011namespace Microsoft . ML . Models
1112{
@@ -18,41 +19,50 @@ private BinaryClassificationMetrics()
1819 {
1920 }
2021
21- internal static BinaryClassificationMetrics FromMetrics ( IHostEnvironment env , IDataView overallMetrics , IDataView confusionMatrix )
22+ internal static List < BinaryClassificationMetrics > FromMetrics ( IHostEnvironment env , IDataView overallMetrics , IDataView confusionMatrix , int confusionMatriceStartIndex = 0 )
2223 {
2324 Contracts . AssertValue ( env ) ;
2425 env . AssertValue ( overallMetrics ) ;
2526 env . AssertValue ( confusionMatrix ) ;
2627
2728 var metricsEnumerable = overallMetrics . AsEnumerable < SerializationClass > ( env , true , ignoreMissingColumns : true ) ;
28- var enumerator = metricsEnumerable . GetEnumerator ( ) ;
29- if ( ! enumerator . MoveNext ( ) )
29+ if ( ! metricsEnumerable . GetEnumerator ( ) . MoveNext ( ) )
3030 {
3131 throw env . Except ( "The overall RegressionMetrics didn't have any rows." ) ;
3232 }
3333
34- SerializationClass metrics = enumerator . Current ;
34+ List < BinaryClassificationMetrics > metrics = new List < BinaryClassificationMetrics > ( ) ;
35+ var confusionMatrices = ConfusionMatrix . Create ( env , confusionMatrix ) . GetEnumerator ( ) ;
3536
36- if ( enumerator . MoveNext ( ) )
37+ int Index = 0 ;
38+ foreach ( var metric in metricsEnumerable )
3739 {
38- throw env . Except ( "The overall RegressionMetrics contained more than 1 row." ) ;
40+
41+ if ( Index ++ >= confusionMatriceStartIndex && ! confusionMatrices . MoveNext ( ) )
42+ {
43+ throw env . Except ( "Confusion matrices didn't have enough matrices." ) ;
44+ }
45+
46+ metrics . Add (
47+ new BinaryClassificationMetrics ( )
48+ {
49+ Auc = metric . Auc ,
50+ Accuracy = metric . Accuracy ,
51+ PositivePrecision = metric . PositivePrecision ,
52+ PositiveRecall = metric . PositiveRecall ,
53+ NegativePrecision = metric . NegativePrecision ,
54+ NegativeRecall = metric . NegativeRecall ,
55+ LogLoss = metric . LogLoss ,
56+ LogLossReduction = metric . LogLossReduction ,
57+ Entropy = metric . Entropy ,
58+ F1Score = metric . F1Score ,
59+ Auprc = metric . Auprc ,
60+ ConfusionMatrix = confusionMatrices . Current ,
61+ } ) ;
62+
3963 }
4064
41- return new BinaryClassificationMetrics ( )
42- {
43- Auc = metrics . Auc ,
44- Accuracy = metrics . Accuracy ,
45- PositivePrecision = metrics . PositivePrecision ,
46- PositiveRecall = metrics . PositiveRecall ,
47- NegativePrecision = metrics . NegativePrecision ,
48- NegativeRecall = metrics . NegativeRecall ,
49- LogLoss = metrics . LogLoss ,
50- LogLossReduction = metrics . LogLossReduction ,
51- Entropy = metrics . Entropy ,
52- F1Score = metrics . F1Score ,
53- Auprc = metrics . Auprc ,
54- ConfusionMatrix = ConfusionMatrix . Create ( env , confusionMatrix ) ,
55- } ;
65+ return metrics ;
5666 }
5767
5868 /// <summary>
@@ -155,7 +165,7 @@ internal static BinaryClassificationMetrics FromMetrics(IHostEnvironment env, ID
155165 /// <summary>
156166 /// This class contains the public fields necessary to deserialize from IDataView.
157167 /// </summary>
158- private class SerializationClass
168+ private sealed class SerializationClass
159169 {
160170#pragma warning disable 649 // never assigned
161171 [ ColumnName ( BinaryClassifierEvaluator . Auc ) ]
0 commit comments