ML.Net Averaged Perceptron only ever returns false

Viewed 78

I’m having a bit of trouble getting my Averaged Perceptron binary classifier to work in ML.Net. I have a simple set of data where if stat1 is greater than or equal to 50 then the result is true, if not false. This training file has around 320 entries.

Example

Stat1 Result
60 TRUE
84 TRUE
69 TRUE
95 TRUE
48 FALSE
88 TRUE
35 FALSE

I’m loading this into my dataview and setting up the label and features columns. The metrics that are retuned indicate that it is accurate (probably too accurate at 100% success rate) but when I test it, it only ever returns false no matter what I put as the input.

Here is my code

static class SimpleTest
{
    public static void RunSimpleTest()
    {
        MLContext mlContext = new MLContext(seed: 0);

        List<TextLoader.Column> mlCols = new List<TextLoader.Column>();

        mlCols.Add(new TextLoader.Column("Stat1", DataKind.Single, 0));
        mlCols.Add(new TextLoader.Column("Result", DataKind.Boolean, 1));

        IDataView dataView = mlContext.Data.LoadFromTextFile("BC_AP CSV Data Simple.csv", mlCols.ToArray(), ',', true, true, true, false);

        var split = mlContext.Data.TrainTestSplit(dataView);

        IDataView trainingDataView = split.TrainSet;
        IDataView testingDataView = split.TestSet;

        IEstimator<ITransformer> pipeline = mlContext.Transforms.CopyColumns("Label", "Result");

        List<string> concatCols = new List<string>()
        {
            "Stat1"
        };

        pipeline = pipeline.Append(mlContext.Transforms.Concatenate("Features", concatCols.ToArray()));
        pipeline = pipeline.Append(mlContext.Transforms.NormalizeMeanVariance("Features", "Features"));

        var trainer = mlContext.BinaryClassification.Trainers.AveragedPerceptron();

        var trainingPipeline = pipeline.Append(trainer);
        var trainedModel = trainingPipeline.Fit(trainingDataView);

        IDataView predictions = trainedModel.Transform(testingDataView);
        var metrics = mlContext.BinaryClassification.EvaluateNonCalibrated(predictions);

        Console.WriteLine($"*Metrics for {trainer.ToString()} classifier model");
        Console.WriteLine(string.Empty);
        Console.WriteLine($"Accuracy: {metrics.Accuracy:F2}");
        Console.WriteLine($"AUC: {metrics.AreaUnderRocCurve:F2}");
        Console.WriteLine($"F1 Score: {metrics.F1Score:F2}");
        Console.WriteLine($"Negative Precision: " + $"{metrics.NegativePrecision:F2}");

        Console.WriteLine($"Negative Recall: {metrics.NegativeRecall:F2}");
        Console.WriteLine($"Positive Precision: " + $"{metrics.PositivePrecision:F2}");

        Console.WriteLine($"Positive Recall: {metrics.PositiveRecall:F2}\n");
        Console.WriteLine(metrics.ConfusionMatrix.GetFormattedConfusionTable());
        Console.WriteLine(string.Empty);

        var predEngine = mlContext.Model.CreatePredictionEngine<SimpleDataObj, SimplePredDataObj>(trainedModel);

        SimpleDataObj dataObj = new SimpleDataObj(99);

        Console.WriteLine("Data for prediction - Expecting True");
        Console.WriteLine(dataObj.ToString());
        Console.WriteLine(string.Empty);

        var predData = predEngine.Predict(dataObj);

        Console.WriteLine("Result");
        Console.WriteLine(predData.Result);
    }

    public class SimpleDataObj
    {
        [LoadColumn(0)]
        public float Stat1 = 0;
        [LoadColumn(1)]
        public bool Result;

        public SimpleDataObj()
        {
        }

        public SimpleDataObj(float stat1)
        {
            Stat1 = stat1;
        }

        public override string ToString()
        {
            StringBuilder sb = new StringBuilder();

            sb.AppendLine("Stat1 = " + Stat1);

            return sb.ToString();
        }
    }

    public class SimplePredDataObj
    {
        [ColumnName("Result")]
        public bool Result { get; set; }
    }
}

I have also created a sample application that will demonstrate the issue in github

https://github.com/XactaAndy/AvgPerceptTest

Any ideas on why it is going wrong?

0 Answers
Related