Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
73 changes: 70 additions & 3 deletions src/Microsoft.ML.Data/Prediction/Calibrator.cs
Original file line number Diff line number Diff line change
Expand Up @@ -1158,7 +1158,7 @@ ICalibrator ICalibratorTrainer.FinishTraining(IChannel ch)
/// <summary>
/// The naive binning-based calibrator.
/// </summary>
public sealed class NaiveCalibrator : ICalibrator, ICanSaveInBinaryFormat
public sealed class NaiveCalibrator : ICalibrator, ICanSaveInBinaryFormat, ISingleCanSaveOnnx
{
internal const string LoaderSignature = "NaiveCaliExec";
internal const string RegistrationName = "NaiveCalibrator";
Expand All @@ -1174,6 +1174,12 @@ private static VersionInfo GetVersionInfo()
loaderAssemblyName: typeof(NaiveCalibrator).Assembly.FullName);
}

/// <summary>
/// Bool required by the interface ISingleCanSaveOnnx, returns true if
/// and only if calibrator can be exported in ONNX.
/// </summary>
bool ICanSaveOnnx.CanSaveOnnx(OnnxContext ctx) => true;

private readonly IHost _host;

/// <summary> The bin size.</summary>
Expand Down Expand Up @@ -1280,6 +1286,48 @@ internal static int GetBinIdx(float output, float min, float binSize, int numBin
return binIdx;
}

bool ISingleCanSaveOnnx.SaveAsOnnx(OnnxContext ctx, string[] outputNames, string featureColumn)
{
_host.CheckValue(ctx, nameof(ctx));
_host.CheckValue(outputNames, nameof(outputNames));
_host.Check(Utils.Size(outputNames) == 2);

const int minimumOpSetVersion = 9;
ctx.CheckOpSetVersion(minimumOpSetVersion, "NaiveCalibrator");

var binProbabilities = ctx.AddInitializer(_binProbs, new long[] { _binProbs.Length, 1 }, "binProbabilities");

string opType = "Sub";
var minVar = ctx.AddInitializer((float)(Min), "Min");
var subNodeOutput = ctx.AddIntermediateVariable(null, "subNodeOutput", true);
var node = ctx.CreateNode(opType, new[] { outputNames[0], minVar }, new[] { subNodeOutput }, ctx.GetNodeName(opType), "");

Comment thread
mstfbl marked this conversation as resolved.
opType = "Div";
var binSizeVar = ctx.AddInitializer((float)(BinSize), "BinSize");
var binIndexOutput = ctx.AddIntermediateVariable(NumberDataViewType.Int32, "binIndexOutput", true);
node = ctx.CreateNode(opType, new[] { subNodeOutput, binSizeVar }, new[] { binIndexOutput }, ctx.GetNodeName(opType), "");

opType = "Cast";
var castOutput = ctx.AddIntermediateVariable(BooleanDataViewType.Instance, "CastOutput");
var castNode = ctx.CreateNode(opType, binIndexOutput, castOutput, ctx.GetNodeName(opType), "");
var t = InternalDataKindExtensions.ToInternalDataKind(DataKind.Boolean).ToType();
castNode.AddAttribute("to", t);

opType = "Not";
var notOutput = ctx.AddIntermediateVariable(BooleanDataViewType.Instance, "IsBinIndexZero");
ctx.CreateNode(opType, castOutput, notOutput, ctx.GetNodeName(opType), "");

Comment thread
mstfbl marked this conversation as resolved.
opType = "Cast";
var castIsBinIndexToInt = ctx.AddIntermediateVariable(NumberDataViewType.Int32, "IsBinIndexAsInt");
var castIsBinIndexToIntNode = ctx.CreateNode(opType, notOutput, castIsBinIndexToInt, ctx.GetNodeName(opType), "");
var t1 = InternalDataKindExtensions.ToInternalDataKind(DataKind.Int32).ToType();
castIsBinIndexToIntNode.AddAttribute("to", t1);

var numBinsVar = ctx.AddInitializer((int)(BinSize), "NumBins");

// TO DO: Complete ONNX conversion.
return true;
}
}

/// <summary>
Expand Down Expand Up @@ -1879,7 +1927,7 @@ public override ICalibrator CreateCalibrator(IChannel ch)
/// <item><description><see cref="Values"/>[n], if x &gt; <see cref="Maxes"/>[n]</description></item>
///</list>
/// </remarks>
public sealed class IsotonicCalibrator : ICalibrator, ICanSaveInBinaryFormat
public sealed class IsotonicCalibrator : ICalibrator, ICanSaveInBinaryFormat, ISingleCanSaveOnnx
{
internal const string LoaderSignature = "PAVCaliExec";
internal const string RegistrationName = "PAVCalibrator";
Expand Down Expand Up @@ -1914,6 +1962,12 @@ private static VersionInfo GetVersionInfo()
/// </summary>
public readonly ImmutableArray<float> Values;

/// <summary>
/// Bool required by the interface ISingleCanSaveOnnx, returns true if
/// and only if calibrator can be exported in ONNX.
/// </summary>
bool ICanSaveOnnx.CanSaveOnnx(OnnxContext ctx) => true;

/// <summary>
/// Initializes a new instance of <see cref="IsotonicCalibrator"/>.
/// </summary>
Expand Down Expand Up @@ -2070,7 +2124,20 @@ private float FindValue(float score)
float t = (score - Maxes[pos - 1]) / (Mins[pos] - Maxes[pos - 1]);
return Values[pos - 1] + t * (Values[pos] - Values[pos - 1]);
}
}

bool ISingleCanSaveOnnx.SaveAsOnnx(OnnxContext ctx, string[] outputNames, string featureColumn)
{
_host.CheckValue(ctx, nameof(ctx));
_host.CheckValue(outputNames, nameof(outputNames));
_host.Check(Utils.Size(outputNames) == 2);

const int minimumOpSetVersion = 9;
ctx.CheckOpSetVersion(minimumOpSetVersion, "IsotonicCalibrator");

// TO DO: Complete ONNX conversion.
return true;
}
}

internal static class Calibrate
{
Expand Down
171 changes: 138 additions & 33 deletions test/Microsoft.ML.Tests/OnnxConversionTest.cs
Original file line number Diff line number Diff line change
Expand Up @@ -261,8 +261,7 @@ public void TestVectorWhiteningOnnxConversionTest()
Done();
}

[Fact]
public void PlattCalibratorOnnxConversionTest()
private (MLContext, IDataView, List<IEstimator<ITransformer>>, EstimatorChain<NormalizingTransformer>) GetEstimatorsForOnnxConversionTests()
Comment thread
mstfbl marked this conversation as resolved.
Outdated
{
var mlContext = new MLContext(seed: 1);
string dataPath = GetDataPath("breast-cancer.txt");
Expand All @@ -289,70 +288,176 @@ public void PlattCalibratorOnnxConversionTest()

var initialPipeline = mlContext.Transforms.ReplaceMissingValues("Features").
Append(mlContext.Transforms.NormalizeMinMax("Features"));
return (mlContext, dataView, estimators, initialPipeline);
}

[Fact]
public void PlattCalibratorOnnxConversionTest()
{
// Step 1: Test calibrator with binary prediction trainer
var (mlContext, dataView, estimators, initialPipeline) = GetEstimatorsForOnnxConversionTests();
foreach (var estimator in estimators)
{
var pipeline = initialPipeline.Append(estimator).Append(mlContext.BinaryClassification.Calibrators.Platt());
var onnxFileName = $"{estimator}-WithPlattCalibrator.onnx";

TestPipeline(pipeline, dataView, onnxFileName, new ColumnComparison[] { new ColumnComparison("Score", 3), new ColumnComparison("PredictedLabel"), new ColumnComparison("Probability", 3) });
}
Done();
}

class PlattModelInput
{
public bool Label { get; set; }
public float Score { get; set; }
}
// Step 2: Test calibrator without any binary prediction trainer
Comment thread
mstfbl marked this conversation as resolved.
IDataView dataSoloCalibrator = mlContext.Data.LoadFromEnumerable(GetCalibratorTestData());

class PlattModelInput2
{
public bool Label { get; set; }
public float ScoreX { get; set; }
var pipelineSoloCalibrator = mlContext.BinaryClassification.Calibrators
.Platt();
var onnxFileNameSoloCalibrator = $"{pipelineSoloCalibrator}-WithPlattCalibrator-Solo.onnx";

TestPipeline(pipelineSoloCalibrator, dataSoloCalibrator,onnxFileNameSoloCalibrator, new ColumnComparison[] { new ColumnComparison("Probability", 3) });

// Step 3: Test calibrator with a non-default Score column name and without any binary prediction trainer
Comment thread
mstfbl marked this conversation as resolved.
IDataView dataSoloCalibratorNonStandard = mlContext.Data.LoadFromEnumerable(GetCalibratorTestDataNonStandard());

var pipelineSoloCalibratorNonStandard = mlContext.BinaryClassification.Calibrators
.Platt(scoreColumnName: "ScoreX");
var onnxFileNameSoloCalibratorNonStandard = $"{pipelineSoloCalibratorNonStandard}-WithPlattCalibrator-Solo-NonStandard.onnx";

TestPipeline(pipelineSoloCalibratorNonStandard, dataSoloCalibratorNonStandard, onnxFileNameSoloCalibratorNonStandard, new ColumnComparison[] { new ColumnComparison("Probability", 3) });

Done();
}

static IEnumerable<PlattModelInput> PlattGetData()
[Fact]
public void FixedPlattCalibratorOnnxConversionTest()
{
for (int i = 0; i < 100; i++)
// Below, FixedPlattCalibrator is utilized by defining slope and offset in Platt's constructor with sample values.
// Step 1: Test calibrator with binary prediction trainer
var (mlContext, dataView, estimators, initialPipeline) = GetEstimatorsForOnnxConversionTests();
foreach (var estimator in estimators)
{
yield return new PlattModelInput { Score = i, Label = i % 2 == 0 };
var pipeline = initialPipeline.Append(estimator).Append(mlContext.BinaryClassification.Calibrators.Platt(slope: -1f, offset: -0.05f));
var onnxFileName = $"{estimator}-WithFixedPlattCalibrator.onnx";

TestPipeline(pipeline, dataView, onnxFileName, new ColumnComparison[] { new ColumnComparison("Score", 3), new ColumnComparison("PredictedLabel"), new ColumnComparison("Probability", 3) });
}

// Step 2: Test calibrator without any binary prediction trainer
IDataView dataSoloCalibrator = mlContext.Data.LoadFromEnumerable(GetCalibratorTestData());

var pipelineSoloCalibrator = mlContext.BinaryClassification.Calibrators
.Platt(slope: -1f, offset: -0.05f);
var onnxFileNameSoloCalibrator = $"{pipelineSoloCalibrator}-WithFixedPlattCalibrator-Solo.onnx";

TestPipeline(pipelineSoloCalibrator, dataSoloCalibrator, onnxFileNameSoloCalibrator, new ColumnComparison[] { new ColumnComparison("Probability", 3) });

// Step 3: Test calibrator with a non-default Score column name and without any binary prediction trainer
IDataView dataSoloCalibratorNonStandard = mlContext.Data.LoadFromEnumerable(GetCalibratorTestDataNonStandard());

var pipelineSoloCalibratorNonStandard = mlContext.BinaryClassification.Calibrators
.Platt(scoreColumnName: "ScoreX", slope: -1f, offset: -0.05f);
var onnxFileNameSoloCalibratorNonStandard = $"{pipelineSoloCalibratorNonStandard}-WithFixedPlattCalibrator-Solo-NonStandard.onnx";

TestPipeline(pipelineSoloCalibratorNonStandard, dataSoloCalibratorNonStandard, onnxFileNameSoloCalibratorNonStandard, new ColumnComparison[] { new ColumnComparison("Probability", 3) });

Done();
}

static IEnumerable<PlattModelInput2> PlattGetData2()
[Fact]
[Trait("Category", "SkipInCI")]
public void NaiveCalibratorOnnxConversionTest()
{
for (int i = 0; i < 100; i++)
// Step 1: Test calibrator with binary prediction trainer
var (mlContext, dataView, estimators, initialPipeline) = GetEstimatorsForOnnxConversionTests();
foreach (var estimator in estimators)
{
yield return new PlattModelInput2 { ScoreX = i, Label = i % 2 == 0 };
var pipeline = initialPipeline.Append(estimator).Append(mlContext.BinaryClassification.Calibrators.Naive());
var onnxFileName = $"{estimator}-WithNaiveCalibrator.onnx";

TestPipeline(pipeline, dataView, onnxFileName, new ColumnComparison[] { new ColumnComparison("Score", 3), new ColumnComparison("PredictedLabel"), new ColumnComparison("Probability", 3) });
}

// Step 2: Test calibrator without any binary prediction trainer
IDataView dataSoloCalibrator = mlContext.Data.LoadFromEnumerable(GetCalibratorTestData());

var pipelineSoloCalibrator = mlContext.BinaryClassification.Calibrators
.Naive();
var onnxFileNameSoloCalibrator = $"{pipelineSoloCalibrator}-WithNaiveCalibrator-Solo.onnx";

TestPipeline(pipelineSoloCalibrator, dataSoloCalibrator, onnxFileNameSoloCalibrator, new ColumnComparison[] { new ColumnComparison("Probability", 3) });

// Step 3: Test calibrator with a non-default Score column name and without any binary prediction trainer
IDataView dataSoloCalibratorNonStandard = mlContext.Data.LoadFromEnumerable(GetCalibratorTestDataNonStandard());

var pipelineSoloCalibratorNonStandard = mlContext.BinaryClassification.Calibrators
.Naive(scoreColumnName: "ScoreX");
var onnxFileNameSoloCalibratorNonStandard = $"{pipelineSoloCalibratorNonStandard}-WithNaiveCalibrator-Solo-NonStandard.onnx";

TestPipeline(pipelineSoloCalibratorNonStandard, dataSoloCalibratorNonStandard, onnxFileNameSoloCalibratorNonStandard, new ColumnComparison[] { new ColumnComparison("Probability", 3) });

Done();
}

[Fact]
public void PlattCalibratorOnnxConversionTest2()
[Trait("Category", "SkipInCI")]
public void IsotonicCalibratorOnnxConversionTest()
{
// Test PlattCalibrator without any binary prediction trainer
var mlContext = new MLContext(seed: 0);
// Step 1: Test calibrator with binary prediction trainer
var (mlContext, dataView, estimators, initialPipeline) = GetEstimatorsForOnnxConversionTests();
foreach (var estimator in estimators)
{
var pipeline = initialPipeline.Append(estimator).Append(mlContext.BinaryClassification.Calibrators.Isotonic());
var onnxFileName = $"{estimator}-WithIsotonicCalibrator.onnx";

IDataView data = mlContext.Data.LoadFromEnumerable(PlattGetData());
TestPipeline(pipeline, dataView, onnxFileName, new ColumnComparison[] { new ColumnComparison("Score", 3), new ColumnComparison("PredictedLabel"), new ColumnComparison("Probability", 3) });
}

var pipeline = mlContext.BinaryClassification.Calibrators
.Platt();
var onnxFileName = $"{pipeline}.onnx";
// Step 2: Test calibrator without any binary prediction trainer
IDataView dataSoloCalibrator = mlContext.Data.LoadFromEnumerable(GetCalibratorTestData());

TestPipeline(pipeline, data, onnxFileName, new ColumnComparison[] { new ColumnComparison("Probability", 3) });
var pipelineSoloCalibrator = mlContext.BinaryClassification.Calibrators
.Isotonic();
var onnxFileNameSoloCalibrator = $"{pipelineSoloCalibrator}-WithIsotonicCalibrator-Solo.onnx";

// Test PlattCalibrator with a non-default Score column name, and without any binary prediction trainer
IDataView data2 = mlContext.Data.LoadFromEnumerable(PlattGetData2());
TestPipeline(pipelineSoloCalibrator, dataSoloCalibrator, onnxFileNameSoloCalibrator, new ColumnComparison[] { new ColumnComparison("Probability", 3) });

var pipeline2 = mlContext.BinaryClassification.Calibrators
.Platt(scoreColumnName: "ScoreX");
var onnxFileName2 = $"{pipeline2}.onnx";
// Step 3: Test calibrator with a non-default Score column name and without any binary prediction trainer
IDataView dataSoloCalibratorNonStandard = mlContext.Data.LoadFromEnumerable(GetCalibratorTestDataNonStandard());

var pipelineSoloCalibratorNonStandard = mlContext.BinaryClassification.Calibrators
.Isotonic(scoreColumnName: "ScoreX");
var onnxFileNameSoloCalibratorNonStandard = $"{pipelineSoloCalibratorNonStandard}-WithIsotonicCalibrator-Solo-NonStandard.onnx";

TestPipeline(pipeline2, data2, onnxFileName2, new ColumnComparison[] { new ColumnComparison("Probability", 3) });
TestPipeline(pipelineSoloCalibratorNonStandard, dataSoloCalibratorNonStandard, onnxFileNameSoloCalibratorNonStandard, new ColumnComparison[] { new ColumnComparison("Probability", 3) });

Done();
}

class ModelInput
Comment thread
mstfbl marked this conversation as resolved.
Outdated
{
Comment thread
mstfbl marked this conversation as resolved.
public bool Label { get; set; }
public float Score { get; set; }
}

class ModelInputNonStandard
{
public bool Label { get; set; }
public float ScoreX { get; set; }
}

static IEnumerable<ModelInput> GetCalibratorTestData()
{
for (int i = 0; i < 100; i++)
{
yield return new ModelInput { Score = i, Label = i % 2 == 0 };
}
}

static IEnumerable<ModelInputNonStandard> GetCalibratorTestDataNonStandard()
{
for (int i = 0; i < 100; i++)
{
yield return new ModelInputNonStandard { ScoreX = i, Label = i % 2 == 0 };
}
}

[Fact]
public void TextNormalizingOnnxConversionTest()
{
Expand Down