diff --git a/src/Microsoft.ML.FastTree/FastTreeClassification.cs b/src/Microsoft.ML.FastTree/FastTreeClassification.cs index 78bcfdd914..3788da21a3 100644 --- a/src/Microsoft.ML.FastTree/FastTreeClassification.cs +++ b/src/Microsoft.ML.FastTree/FastTreeClassification.cs @@ -90,6 +90,8 @@ internal static IPredictorProducing Create(IHostEnvironment env, ModelLoa env.CheckValue(ctx, nameof(ctx)); ctx.CheckAtModel(GetVersionInfo()); var predictor = new FastTreeBinaryModelParameters(env, ctx); + // Preserve compatibility with archives that embed a calibrator inside the predictor. + // Current calibrated models save the predictor and calibrator as siblings in a separate wrapper. ICalibrator calibrator; ctx.LoadModelOrNull(env, out calibrator, @"Calibrator"); if (calibrator == null) diff --git a/src/Microsoft.ML.FastTree/GamClassification.cs b/src/Microsoft.ML.FastTree/GamClassification.cs index 9f2ff943cc..d34fac1ecd 100644 --- a/src/Microsoft.ML.FastTree/GamClassification.cs +++ b/src/Microsoft.ML.FastTree/GamClassification.cs @@ -238,6 +238,8 @@ internal static IPredictorProducing Create(IHostEnvironment env, ModelLoa ctx.CheckAtModel(GetVersionInfo()); var predictor = new GamBinaryModelParameters(env, ctx); + // Preserve compatibility with archives that embed a calibrator inside the predictor. + // Current calibrated models save the predictor and calibrator as siblings in a separate wrapper. ICalibrator calibrator; ctx.LoadModelOrNull(env, out calibrator, @"Calibrator"); if (calibrator == null) diff --git a/src/Microsoft.ML.FastTree/RandomForestClassification.cs b/src/Microsoft.ML.FastTree/RandomForestClassification.cs index a63811576e..4ee6223ed8 100644 --- a/src/Microsoft.ML.FastTree/RandomForestClassification.cs +++ b/src/Microsoft.ML.FastTree/RandomForestClassification.cs @@ -109,6 +109,8 @@ internal static IPredictorProducing Create(IHostEnvironment env, ModelLoa env.CheckValue(ctx, nameof(ctx)); ctx.CheckAtModel(GetVersionInfo()); var predictor = new FastForestBinaryModelParameters(env, ctx); + // Preserve compatibility with archives that embed a calibrator inside the predictor. + // Current calibrated models save the predictor and calibrator as siblings in a separate wrapper. ICalibrator calibrator; ctx.LoadModelOrNull(env, out calibrator, @"Calibrator"); if (calibrator == null) diff --git a/src/Microsoft.ML.LightGbm/LightGbmBinaryTrainer.cs b/src/Microsoft.ML.LightGbm/LightGbmBinaryTrainer.cs index c6d56aae2f..094a3382f4 100644 --- a/src/Microsoft.ML.LightGbm/LightGbmBinaryTrainer.cs +++ b/src/Microsoft.ML.LightGbm/LightGbmBinaryTrainer.cs @@ -78,6 +78,8 @@ internal static IPredictorProducing Create(IHostEnvironment env, ModelLoa env.CheckValue(ctx, nameof(ctx)); ctx.CheckAtModel(GetVersionInfo()); var predictor = new LightGbmBinaryModelParameters(env, ctx); + // Preserve compatibility with archives that embed a calibrator inside the predictor. + // Current calibrated models save the predictor and calibrator as siblings in a separate wrapper. ICalibrator calibrator; ctx.LoadModelOrNull(env, out calibrator, @"Calibrator"); if (calibrator == null) diff --git a/src/Microsoft.ML.StandardTrainers/Standard/LinearModelParameters.cs b/src/Microsoft.ML.StandardTrainers/Standard/LinearModelParameters.cs index 4093bbf571..49aee8effd 100644 --- a/src/Microsoft.ML.StandardTrainers/Standard/LinearModelParameters.cs +++ b/src/Microsoft.ML.StandardTrainers/Standard/LinearModelParameters.cs @@ -484,6 +484,8 @@ internal static IPredictorProducing Create(IHostEnvironment env, ModelLoa env.CheckValue(ctx, nameof(ctx)); ctx.CheckAtModel(GetVersionInfo()); var predictor = new LinearBinaryModelParameters(env, ctx); + // Preserve compatibility with archives that embed a calibrator inside the predictor. + // Current calibrated models save the predictor and calibrator as siblings in a separate wrapper. ICalibrator calibrator; ctx.LoadModelOrNull(env, out calibrator, @"Calibrator"); if (calibrator == null) diff --git a/test/Microsoft.ML.Tests/CalibratedModelParametersTests.cs b/test/Microsoft.ML.Tests/CalibratedModelParametersTests.cs index 0b0e4b241f..573e6b3564 100644 --- a/test/Microsoft.ML.Tests/CalibratedModelParametersTests.cs +++ b/test/Microsoft.ML.Tests/CalibratedModelParametersTests.cs @@ -3,12 +3,15 @@ // See the LICENSE file in the project root for more information. using System; +using System.IO; using Microsoft.ML.Calibrators; using Microsoft.ML.Data; using Microsoft.ML.Internal.Utilities; +using Microsoft.ML.Model; using Microsoft.ML.RunTests; using Microsoft.ML.Trainers; using Microsoft.ML.Trainers.FastTree; +using Microsoft.ML.Trainers.LightGbm; using Xunit; using Xunit.Abstractions; @@ -86,7 +89,230 @@ public void TestFeatureWeightsCalibratedModelParametersLoading() Done(); } - #region Helpers + [Theory] + [InlineData(typeof(LinearBinaryModelParameters))] + [InlineData(typeof(GamBinaryModelParameters))] + [InlineData(typeof(FastTreeBinaryModelParameters))] + [InlineData(typeof(FastForestBinaryModelParameters))] + [InlineData(typeof(LightGbmBinaryModelParameters))] + public void BinaryModelWithoutEmbeddedCalibratorLoadsUnwrapped(Type modelType) + { + var loaded = RoundTripBinaryModel(CreateBinaryModel(modelType)); + + Assert.Equal(modelType, loaded.GetType()); + Assert.False(loaded is CalibratedModelParametersBase); + AssertBinaryScores(loaded); + Done(); + } + + [Theory] + [InlineData(typeof(LinearBinaryModelParameters), false)] + [InlineData(typeof(LinearBinaryModelParameters), true)] + [InlineData(typeof(GamBinaryModelParameters), false)] + [InlineData(typeof(FastTreeBinaryModelParameters), false)] + [InlineData(typeof(FastForestBinaryModelParameters), false)] + [InlineData(typeof(LightGbmBinaryModelParameters), false)] + public void EmbeddedCalibratorLoadsAndResavesInWrapperFormat(Type modelType, bool useNaiveCalibrator) + { + ICalibrator calibrator = useNaiveCalibrator + ? new NaiveCalibrator(Env, min: -4, binSize: 4, binProbs: new[] { 0.125f, 0.875f }) + : new PlattCalibrator(Env, slope: -0.75, offset: 0.25); + Type wrapperType = modelType == typeof(LinearBinaryModelParameters) && !useNaiveCalibrator + ? typeof(ParameterMixingCalibratedModelParameters<,>) + : modelType == typeof(LightGbmBinaryModelParameters) + ? typeof(ValueMapperCalibratedModelParameters<,>) + : typeof(SchemaBindableCalibratedModelParameters<,>); + + // These synthetic archives exercise the supported layout, not a particular historical release. + var loaded = RoundTripBinaryModel(CreateBinaryModel(modelType), calibrator); + + Assert.Equal(wrapperType.MakeGenericType(modelType, typeof(ICalibrator)), loaded.GetType()); + AssertCalibratedBinaryModel(loaded, modelType, calibrator, wrapperType); + AssertBinaryScores(loaded, calibrated: true, useNaiveCalibrator: useNaiveCalibrator); + + var reloaded = RoundTripBinaryModel(loaded); + + AssertCalibratedBinaryModel(reloaded, modelType, calibrator, wrapperType); + AssertBinaryScores(reloaded, calibrated: true, useNaiveCalibrator: useNaiveCalibrator); + Done(); + } + + [Theory] + [InlineData(typeof(LinearBinaryModelParameters))] + [InlineData(typeof(GamBinaryModelParameters))] + [InlineData(typeof(FastTreeBinaryModelParameters))] + [InlineData(typeof(FastForestBinaryModelParameters))] + [InlineData(typeof(LightGbmBinaryModelParameters))] + public void CalibratedBinaryModelRoundTripPreservesWrapperAndPredictions(Type modelType) + { + var calibrator = new PlattCalibrator(Env, slope: -0.75, offset: 0.25); + var model = CreateCalibratedBinaryModel(CreateBinaryModel(modelType), calibrator); + + var loaded = RoundTripBinaryModel(model); + + Assert.Equal(model.GetType(), loaded.GetType()); + AssertCalibratedBinaryModel(loaded, modelType, calibrator, model.GetType().GetGenericTypeDefinition()); + AssertBinaryScores(loaded, calibrated: true); + Done(); + } + + [Theory] + [InlineData(typeof(LinearBinaryModelParameters))] + [InlineData(typeof(GamBinaryModelParameters))] + [InlineData(typeof(FastTreeBinaryModelParameters))] + [InlineData(typeof(FastForestBinaryModelParameters))] + [InlineData(typeof(LightGbmBinaryModelParameters))] + public void InvalidEmbeddedCalibratorFailsToLoad(Type modelType) + { + var model = CreateBinaryModel(modelType); + var calibrator = new PlattCalibrator(Env, slope: double.NaN, offset: 0.25); + + var exception = Assert.Throws(() => RoundTripBinaryModel(model, calibrator)); + + // The component catalogue wraps the calibrator's decoding failure. + Assert.IsType(exception.GetBaseException()); + Done(); + } + + private IPredictorProducing CreateBinaryModel(Type modelType) + { + // All five models score -2.5 for [-1] and 1.5 for [1], without training or native libraries. + if (modelType == typeof(LinearBinaryModelParameters)) + return new LinearBinaryModelParameters(Env, new VBuffer(1, new[] { 2f }), bias: -0.5f); + if (modelType == typeof(GamBinaryModelParameters)) + { + return new GamBinaryModelParameters(Env, + new[] { new[] { 0d, double.PositiveInfinity } }, new[] { new[] { -2d, 2d } }, + intercept: -0.5, inputLength: 1, featureToInputMap: null); + } + + InternalRegressionTree tree = modelType == typeof(FastForestBinaryModelParameters) + ? new InternalQuantileRegressionTree( + splitFeatures: new[] { 0 }, splitGain: new[] { 1d }, gainPValue: null, + rawThresholds: new[] { 0f }, defaultValueForMissing: null, + lteChild: new[] { -1 }, gtChild: new[] { -2 }, leafValues: new[] { -2.5, 1.5 }, + categoricalSplitFeatures: new int[1][], categoricalSplit: new bool[1]) + : new InternalRegressionTree( + splitFeatures: new[] { 0 }, splitGain: new[] { 1d }, gainPValue: null, + rawThresholds: new[] { 0f }, defaultValueForMissing: null, + lteChild: new[] { -1 }, gtChild: new[] { -2 }, leafValues: new[] { -2.5, 1.5 }, + categoricalSplitFeatures: new int[1][], categoricalSplit: new bool[1]); + var ensemble = new InternalTreeEnsemble(); + ensemble.AddTree(tree); + + if (modelType == typeof(FastTreeBinaryModelParameters)) + return new FastTreeBinaryModelParameters(Env, ensemble, featureCount: 1, innerArgs: null); + if (modelType == typeof(FastForestBinaryModelParameters)) + return new FastForestBinaryModelParameters(Env, ensemble, featureCount: 1, innerArgs: null); + if (modelType == typeof(LightGbmBinaryModelParameters)) + return new LightGbmBinaryModelParameters(Env, ensemble, featureCount: 1, innerArgs: null); + throw new ArgumentOutOfRangeException(nameof(modelType)); + } + + private IPredictorProducing CreateCalibratedBinaryModel(IPredictorProducing model, PlattCalibrator calibrator) + { + return model switch + { + LinearBinaryModelParameters linear => new ParameterMixingCalibratedModelParameters(Env, linear, calibrator), + GamBinaryModelParameters gam => new ValueMapperCalibratedModelParameters(Env, gam, calibrator), + FastTreeBinaryModelParameters fastTree => new FeatureWeightsCalibratedModelParameters(Env, fastTree, calibrator), + FastForestBinaryModelParameters fastForest => new FeatureWeightsCalibratedModelParameters(Env, fastForest, calibrator), + LightGbmBinaryModelParameters lightGbm => new FeatureWeightsCalibratedModelParameters(Env, lightGbm, calibrator), + _ => throw new ArgumentOutOfRangeException(nameof(model)) + }; + } + + private IPredictorProducing RoundTripBinaryModel(IPredictorProducing model, ICalibrator embeddedCalibrator = null) + { + using var stream = new MemoryStream(); + using (var writer = RepositoryWriter.CreateNew(stream, Env, useFileSystem: false)) + { + ModelSaveContext.SaveModel(writer, model, "Predictor"); + if (embeddedCalibrator != null) + ModelSaveContext.SaveModel(writer, embeddedCalibrator, Path.Combine("Predictor", "Calibrator")); + writer.Commit(); + } + + stream.Position = 0; + using var reader = RepositoryReader.Open(stream, Env, useFileSystem: false); + if (model is CalibratedModelParametersBase calibrated) + { + // Current-format archives have a bare predictor and its calibrator as siblings. + ModelLoadContext.LoadModel, SignatureLoadModel>( + Env, out var subModel, reader, Path.Combine("Predictor", "Predictor")); + Assert.Equal(calibrated.SubModel.GetType(), subModel.GetType()); + ModelLoadContext.LoadModel( + Env, out var calibrator, reader, Path.Combine("Predictor", "Calibrator")); + Assert.Equal(calibrated.Calibrator.GetType(), calibrator.GetType()); + } + + ModelLoadContext.LoadModel, SignatureLoadModel>(Env, out var loaded, reader, "Predictor"); + return loaded; + } + + private static void AssertCalibratedBinaryModel(IPredictorProducing model, Type modelType, + ICalibrator expectedCalibrator, Type wrapperType) + { + var calibrated = Assert.IsAssignableFrom(model); + Assert.Equal(modelType, calibrated.SubModel.GetType()); + Assert.Equal(wrapperType, model.GetType().GetGenericTypeDefinition()); + Assert.IsAssignableFrom>(model); + Assert.Equal(wrapperType == typeof(ParameterMixingCalibratedModelParameters<,>), model is IParameterMixer); + Assert.Equal(wrapperType != typeof(SchemaBindableCalibratedModelParameters<,>), model is IValueMapperDist); + if (wrapperType == typeof(SchemaBindableCalibratedModelParameters<,>)) + Assert.IsAssignableFrom(model); + + if (expectedCalibrator is PlattCalibrator platt) + { + var actual = Assert.IsType(calibrated.Calibrator); + Assert.Equal(platt.Slope, actual.Slope); + Assert.Equal(platt.Offset, actual.Offset); + } + else + { + var naive = Assert.IsType(expectedCalibrator); + var actual = Assert.IsType(calibrated.Calibrator); + Assert.Equal(naive.Min, actual.Min); + Assert.Equal(naive.BinSize, actual.BinSize); + Assert.Equal(naive.BinProbs, actual.BinProbs); + } + } + + private void AssertBinaryScores(IPredictorProducing model, bool calibrated = false, bool useNaiveCalibrator = false) + { + var builder = new ArrayDataViewBuilder(Env); + builder.AddColumn("Features", NumberDataViewType.Single, new[] { -1f }, new[] { 1f }); + var data = builder.GetDataView(); + var bindable = ScoreUtils.GetSchemaBindableMapper(Env, model); + var mapper = Assert.IsAssignableFrom( + bindable.Bind(Env, new RoleMappedSchema(data.Schema, label: null, feature: "Features"))); + Assert.Equal(calibrated, mapper.OutputSchema.TryGetColumnIndex("Probability", out int probabilityColumn)); + + using var cursor = data.GetRowCursor(data.Schema); + using var row = mapper.GetRow(cursor, mapper.OutputSchema); + var scoreGetter = row.GetGetter(row.Schema["Score"]); + var probabilityGetter = calibrated ? row.GetGetter(row.Schema[probabilityColumn]) : null; + var expectedScores = new[] { -2.5f, 1.5f }; + var expectedProbabilities = useNaiveCalibrator + ? new[] { 0.125f, 0.875f } + : new[] { (float)(1 / (1 + Math.Exp(2.125))), (float)(1 / (1 + Math.Exp(-0.875))) }; + + for (int i = 0; i < expectedScores.Length; i++) + { + Assert.True(cursor.MoveNext()); + float score = 0; + scoreGetter(ref score); + Assert.Equal(expectedScores[i], score); + if (calibrated) + { + float probability = 0; + probabilityGetter(ref probability); + Assert.Equal(expectedProbabilities[i], probability, precision: 6); + } + } + Assert.False(cursor.MoveNext()); + } + /// /// Features: x1, x2, x3, xRand; y = 10*x1 + 20x2 + 5.5x3 + e, xRand- random and Label y is to dependant on xRand. /// xRand has the least importance: Evaluation metrics do not change a lot when xRand is permuted. @@ -156,6 +382,5 @@ private float GetArrayAverage(float[] scores) return averageScore; } - #endregion } }