From f5b09b1c648582b09cc4e147e9ff0836952ab1bc Mon Sep 17 00:00:00 2001 From: XiaoYun Zhang Date: Wed, 25 May 2022 12:52:36 -0700 Subject: [PATCH 1/7] implement auto featurizer --- src/Microsoft.ML.AutoML/API/AutoCatalog.cs | 110 +++++++++++++++--- src/Microsoft.ML.Data/Transforms/Hashing.cs | 2 +- .../Text/NgramHashingTransformer.cs | 2 +- ...ests.AutoFeaturizer_iris_test.approved.txt | 34 ++++++ ...AutoFeaturizer_uci_adult_test.approved.txt | 83 +++++++++++++ .../AutoFeaturizerTests.cs | 65 +++++++++++ test/Microsoft.ML.AutoML.Tests/DatasetUtil.cs | 15 +++ .../GridSearchTunerTests.cs | 2 +- 8 files changed, 294 insertions(+), 19 deletions(-) create mode 100644 test/Microsoft.ML.AutoML.Tests/ApprovalTests/AutoFeaturizerTests.AutoFeaturizer_iris_test.approved.txt create mode 100644 test/Microsoft.ML.AutoML.Tests/ApprovalTests/AutoFeaturizerTests.AutoFeaturizer_uci_adult_test.approved.txt create mode 100644 test/Microsoft.ML.AutoML.Tests/AutoFeaturizerTests.cs diff --git a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs index 6937e4e59e..2b6385d06f 100644 --- a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs +++ b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs @@ -4,8 +4,11 @@ using System; using System.Collections.Generic; +using System.Diagnostics.Contracts; +using System.Linq; using Microsoft.ML.AutoML.CodeGen; using Microsoft.ML.Data; +using Microsoft.ML.Runtime; using Microsoft.ML.SearchSpace; using Microsoft.ML.Trainers.FastTree; @@ -538,32 +541,53 @@ public SweepableEstimator[] Regression(string labelColumnName = DefaultColumnNam /// input column name. internal SweepableEstimator[] TextFeaturizer(string outputColumnName, string inputColumnName) { - throw new NotImplementedException(); + var option = new FeaturizeTextOption + { + InputColumnName = inputColumnName, + OutputColumnName = outputColumnName, + }; + + return new[] { SweepableEstimatorFactory.CreateFeaturizeText(option) }; } /// /// Create a list of for featurizing numeric columns. /// - /// output column name. - /// input column name. - internal SweepableEstimator[] NumericFeaturizer(string outputColumnName, string inputColumnName) + /// output column names. + /// input column names. + internal SweepableEstimator[] NumericFeaturizer(string[] outputColumnNames, string[] inputColumnNames) { - throw new NotImplementedException(); + Contracts.Check(outputColumnNames.Count() == inputColumnNames.Count() && outputColumnNames.Count() > 0, "outputColumnNames and inputColumnNames must have the same length and greater than 0"); + var replaceMissingValueOption = new ReplaceMissingValueOption + { + InputColumnNames = inputColumnNames, + OutputColumnNames = outputColumnNames, + }; + + return new[] { SweepableEstimatorFactory.CreateReplaceMissingValues(replaceMissingValueOption) }; } /// /// Create a list of for featurizing catalog columns. /// - /// output column name. - /// input column name. - internal SweepableEstimator[] CatalogFeaturizer(string outputColumnName, string inputColumnName) + /// output column names. + /// input column names. + internal SweepableEstimator[] CatalogFeaturizer(string[] outputColumnNames, string[] inputColumnNames) { - throw new NotImplementedException(); + Contracts.Check(outputColumnNames.Count() == inputColumnNames.Count() && outputColumnNames.Count() > 0, "outputColumnNames and inputColumnNames must have the same length and greater than 0"); + + var option = new OneHotOption + { + InputColumnNames = inputColumnNames, + OutputColumnNames = outputColumnNames, + }; + + return new SweepableEstimator[] { SweepableEstimatorFactory.CreateOneHotEncoding(option), SweepableEstimatorFactory.CreateOneHotHashEncoding(option) }; } /// /// Create a single featurize pipeline according to . This function will collect all columns in and not in , - /// featurizing them using , or . And combine + /// featurizing them using , or . And combine /// them into a single feature column as output. /// /// input data. @@ -572,21 +596,75 @@ internal SweepableEstimator[] CatalogFeaturizer(string outputColumnName, string /// columns that won't be included when featurizing, like label internal MultiModelPipeline Featurizer(IDataView data, string outputColumnName = "Features", string[] catalogColumns = null, string[] excludeColumns = null) { - throw new NotImplementedException(); + var columnInfo = new ColumnInformation(); + + if (excludeColumns != null) + { + foreach (var ignoreColumn in excludeColumns) + { + columnInfo.IgnoredColumnNames.Add(ignoreColumn); + } + } + + if (catalogColumns != null) + { + foreach (var catalogColumn in catalogColumns) + { + columnInfo.CategoricalColumnNames.Add(catalogColumn); + } + } + + return this.Featurizer(data, columnInfo, outputColumnName); } /// - /// Create a single featurize pipeline according to . This function will collect all columns in and not in , - /// featurizing them using , or . And combine + /// Create a single featurize pipeline according to . This function will collect all columns in , + /// featurizing them using , or . And combine /// them into a single feature column as output. /// + /// input data. /// column information. /// output feature column. - /// columns that won't be included when featurizing, like label /// - internal MultiModelPipeline Featurizer(ColumnInformation columnInformation, string outputColumnName = "Features", string[] excludeColumns = null) + internal MultiModelPipeline Featurizer(IDataView data, ColumnInformation columnInformation, string outputColumnName = "Features") { - throw new NotImplementedException(); + var columnPurposes = PurposeInference.InferPurposes(this._context, data, columnInformation); + var textFeatures = columnPurposes.Where(c => c.Purpose == ColumnPurpose.TextFeature); + var numericFeatures = columnPurposes.Where(c => c.Purpose == ColumnPurpose.NumericFeature); + var catalogFeatures = columnPurposes.Where(c => c.Purpose == ColumnPurpose.CategoricalFeature); + var textFeatureColumnNames = textFeatures.Select(c => data.Schema[c.ColumnIndex].Name).ToArray(); + var numericFeatureColumnNames = numericFeatures.Select(c => data.Schema[c.ColumnIndex].Name).ToArray(); + var catalogFeatureColumnNames = catalogFeatures.Select(c => data.Schema[c.ColumnIndex].Name).ToArray(); + + var pipeline = new MultiModelPipeline(); + if (numericFeatureColumnNames.Length > 0) + { + pipeline = pipeline.Append(this.NumericFeaturizer(numericFeatureColumnNames, numericFeatureColumnNames)); + } + + if (catalogFeatureColumnNames.Length > 0) + { + pipeline = pipeline.Append(this.CatalogFeaturizer(catalogFeatureColumnNames, catalogFeatureColumnNames)); + } + + foreach (var textColumn in textFeatureColumnNames) + { + pipeline = pipeline.Append(this.TextFeaturizer(textColumn, textColumn)); + } + + var option = new ConcatOption + { + InputColumnNames = textFeatureColumnNames.Concat(numericFeatureColumnNames).Concat(catalogFeatureColumnNames).ToArray(), + OutputColumnName = outputColumnName, + }; + + if (option.InputColumnNames.Length > 0) + { + pipeline = pipeline.Append(SweepableEstimatorFactory.CreateConcatenate(option)); + } + + return pipeline; + } } } diff --git a/src/Microsoft.ML.Data/Transforms/Hashing.cs b/src/Microsoft.ML.Data/Transforms/Hashing.cs index 8e7dd10221..d3131ea2c1 100644 --- a/src/Microsoft.ML.Data/Transforms/Hashing.cs +++ b/src/Microsoft.ML.Data/Transforms/Hashing.cs @@ -182,7 +182,7 @@ internal HashingTransformer(IHostEnvironment env, params HashingEstimator.Column foreach (var column in _columns) { if (column.MaximumNumberOfInverts != 0) - throw Host.ExceptParam(nameof(columns), $"Found column with {nameof(column.MaximumNumberOfInverts)} set to non zero value, please use { nameof(HashingEstimator)} instead"); + throw Host.ExceptParam(nameof(columns), $"Found column with {nameof(column.MaximumNumberOfInverts)} set to non zero value, please use {nameof(HashingEstimator)} instead"); if (column.Combine && column.UseOrderedHashing) throw Host.ExceptParam(nameof(HashingEstimator.ColumnOptions.Combine), "When the 'Combine' option is specified, ordered hashing is not supported."); diff --git a/src/Microsoft.ML.Transforms/Text/NgramHashingTransformer.cs b/src/Microsoft.ML.Transforms/Text/NgramHashingTransformer.cs index 8b75eb0d3c..484f6060c2 100644 --- a/src/Microsoft.ML.Transforms/Text/NgramHashingTransformer.cs +++ b/src/Microsoft.ML.Transforms/Text/NgramHashingTransformer.cs @@ -184,7 +184,7 @@ internal NgramHashingTransformer(IHostEnvironment env, params NgramHashingEstima foreach (var column in _columns) { if (column.MaximumNumberOfInverts != 0) - throw Host.ExceptParam(nameof(columns), $"Found colunm with {nameof(column.MaximumNumberOfInverts)} set to non zero value, please use { nameof(NgramHashingEstimator)} instead"); + throw Host.ExceptParam(nameof(columns), $"Found colunm with {nameof(column.MaximumNumberOfInverts)} set to non zero value, please use {nameof(NgramHashingEstimator)} instead"); } } diff --git a/test/Microsoft.ML.AutoML.Tests/ApprovalTests/AutoFeaturizerTests.AutoFeaturizer_iris_test.approved.txt b/test/Microsoft.ML.AutoML.Tests/ApprovalTests/AutoFeaturizerTests.AutoFeaturizer_iris_test.approved.txt new file mode 100644 index 0000000000..e6ec8ed089 --- /dev/null +++ b/test/Microsoft.ML.AutoML.Tests/ApprovalTests/AutoFeaturizerTests.AutoFeaturizer_iris_test.approved.txt @@ -0,0 +1,34 @@ +{ + "schema": "e0 * e1", + "estimators": { + "e0": { + "estimatorType": "ReplaceMissingValues", + "parameter": { + "OutputColumnNames": [ + "col1", + "col2", + "col3", + "col4" + ], + "InputColumnNames": [ + "col1", + "col2", + "col3", + "col4" + ] + } + }, + "e1": { + "estimatorType": "Concatenate", + "parameter": { + "InputColumnNames": [ + "col1", + "col2", + "col3", + "col4" + ], + "OutputColumnName": "Features" + } + } + } +} \ No newline at end of file diff --git a/test/Microsoft.ML.AutoML.Tests/ApprovalTests/AutoFeaturizerTests.AutoFeaturizer_uci_adult_test.approved.txt b/test/Microsoft.ML.AutoML.Tests/ApprovalTests/AutoFeaturizerTests.AutoFeaturizer_uci_adult_test.approved.txt new file mode 100644 index 0000000000..63a64233a4 --- /dev/null +++ b/test/Microsoft.ML.AutoML.Tests/ApprovalTests/AutoFeaturizerTests.AutoFeaturizer_uci_adult_test.approved.txt @@ -0,0 +1,83 @@ +{ + "schema": "e0 * (e1 \u002B e2) * e3", + "estimators": { + "e0": { + "estimatorType": "ReplaceMissingValues", + "parameter": { + "OutputColumnNames": [ + "Features" + ], + "InputColumnNames": [ + "Features" + ] + } + }, + "e1": { + "estimatorType": "OneHotEncoding", + "parameter": { + "OutputColumnNames": [ + "Workclass", + "education", + "marital-status", + "occupation", + "relationship", + "ethnicity", + "sex", + "native-country-region" + ], + "InputColumnNames": [ + "Workclass", + "education", + "marital-status", + "occupation", + "relationship", + "ethnicity", + "sex", + "native-country-region" + ] + } + }, + "e2": { + "estimatorType": "OneHotHashEncoding", + "parameter": { + "OutputColumnNames": [ + "Workclass", + "education", + "marital-status", + "occupation", + "relationship", + "ethnicity", + "sex", + "native-country-region" + ], + "InputColumnNames": [ + "Workclass", + "education", + "marital-status", + "occupation", + "relationship", + "ethnicity", + "sex", + "native-country-region" + ] + } + }, + "e3": { + "estimatorType": "Concatenate", + "parameter": { + "InputColumnNames": [ + "Features", + "Workclass", + "education", + "marital-status", + "occupation", + "relationship", + "ethnicity", + "sex", + "native-country-region" + ], + "OutputColumnName": "OutputFeature" + } + } + } +} \ No newline at end of file diff --git a/test/Microsoft.ML.AutoML.Tests/AutoFeaturizerTests.cs b/test/Microsoft.ML.AutoML.Tests/AutoFeaturizerTests.cs new file mode 100644 index 0000000000..aa66207921 --- /dev/null +++ b/test/Microsoft.ML.AutoML.Tests/AutoFeaturizerTests.cs @@ -0,0 +1,65 @@ +// Licensed to the .NET Foundation under one or more agreements. +// The .NET Foundation licenses this file to you under the MIT license. +// See the LICENSE file in the project root for more information. + +using System; +using System.Collections.Generic; +using System.Text; +using System.Text.Json; +using Microsoft.ML.TestFramework; +using Xunit; +using Xunit.Abstractions; +using ApprovalTests; +using ApprovalTests.Namers; +using ApprovalTests.Reporters; +using System.Text.Json.Serialization; + +namespace Microsoft.ML.AutoML.Test +{ + public class AutoFeaturizerTests : BaseTestClass + { + private readonly JsonSerializerOptions _jsonSerializerOptions; + + public AutoFeaturizerTests(ITestOutputHelper output) + : base(output) + { + _jsonSerializerOptions = new JsonSerializerOptions() + { + WriteIndented = true, + Converters = + { + new JsonStringEnumConverter(), new DoubleToDecimalConverter(), new FloatToDecimalConverter(), + }, + }; + + if (Environment.GetEnvironmentVariable("HELIX_CORRELATION_ID") != null) + { + Approvals.UseAssemblyLocationForApprovedFiles(); + } + } + + [Fact] + [UseReporter(typeof(DiffReporter))] + [UseApprovalSubdirectory("ApprovalTests")] + public void AutoFeaturizer_uci_adult_test() + { + var context = new MLContext(1); + var dataset = DatasetUtil.GetUciAdultDataView(); + var pipeline = context.Auto().Featurizer(dataset, outputColumnName: "OutputFeature", excludeColumns: new[] { "Label" }); + + Approvals.Verify(JsonSerializer.Serialize(pipeline, _jsonSerializerOptions)); + } + + [Fact] + [UseReporter(typeof(DiffReporter))] + [UseApprovalSubdirectory("ApprovalTests")] + public void AutoFeaturizer_iris_test() + { + var context = new MLContext(1); + var dataset = DatasetUtil.GetIrisDataView(); + var pipeline = context.Auto().Featurizer(dataset, excludeColumns: new[] { "Label" }); + + Approvals.Verify(JsonSerializer.Serialize(pipeline, _jsonSerializerOptions)); + } + } +} diff --git a/test/Microsoft.ML.AutoML.Tests/DatasetUtil.cs b/test/Microsoft.ML.AutoML.Tests/DatasetUtil.cs index a465e55fcd..17cdfb2791 100644 --- a/test/Microsoft.ML.AutoML.Tests/DatasetUtil.cs +++ b/test/Microsoft.ML.AutoML.Tests/DatasetUtil.cs @@ -24,6 +24,8 @@ internal static class DatasetUtil private static IDataView _uciAdultDataView; + private static IDataView _irisDataView; + public static string GetUciAdultDataset() => GetDataPath("adult.tiny.with-schema.txt"); public static string GetMlNetGeneratedRegressionDataset() => GetDataPath("generated_regression_dataset.csv"); @@ -50,6 +52,19 @@ public static IDataView GetUciAdultDataView() return _uciAdultDataView; } + public static IDataView GetIrisDataView() + { + if (_irisDataView == null) + { + var context = new MLContext(1); + var dataFile = GetIrisDataset(); + var columnInferenceResult = context.Auto().InferColumns(dataFile, 0, groupColumns: false); + var textLoader = context.Data.CreateTextLoader(columnInferenceResult.TextLoaderOptions); + _irisDataView = textLoader.Load(dataFile); + } + return _irisDataView; + } + public static string GetFlowersDataset() { const string datasetName = @"flowers"; diff --git a/test/Microsoft.ML.AutoML.Tests/GridSearchTunerTests.cs b/test/Microsoft.ML.AutoML.Tests/GridSearchTunerTests.cs index b59aada665..fab4f0de67 100644 --- a/test/Microsoft.ML.AutoML.Tests/GridSearchTunerTests.cs +++ b/test/Microsoft.ML.AutoML.Tests/GridSearchTunerTests.cs @@ -12,7 +12,7 @@ using Xunit; using Xunit.Abstractions; -namespace Microsoft.ML.AutoML.Tests +namespace Microsoft.ML.AutoML.Test { public class GridSearchTunerTests : BaseTestClass { From 564ee43f6ec3868d3afe09c983bbf399a0a6b25c Mon Sep 17 00:00:00 2001 From: XiaoYun Zhang Date: Wed, 25 May 2022 13:16:40 -0700 Subject: [PATCH 2/7] detect overlapping --- src/Microsoft.ML.AutoML/API/AutoCatalog.cs | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs index 2b6385d06f..239bebfbdc 100644 --- a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs +++ b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs @@ -592,10 +592,22 @@ internal SweepableEstimator[] CatalogFeaturizer(string[] outputColumnNames, stri /// /// input data. /// columns that should be treated as catalog. If not specified, it will automatically infer if a column is catalog or not. + /// columns that should be treated as numeric. If not specified, it will automatically infer if a column is catalog or not. + /// columns that should be treated as text. If not specified, it will automatically infer if a column is catalog or not. /// output feature column. /// columns that won't be included when featurizing, like label - internal MultiModelPipeline Featurizer(IDataView data, string outputColumnName = "Features", string[] catalogColumns = null, string[] excludeColumns = null) + public MultiModelPipeline Featurizer(IDataView data, string outputColumnName = "Features", string[] catalogColumns = null, string[] numericColumns = null, string[] textColumns = null, string[] excludeColumns = null) { + // validate if there's overlapping among catalogColumns, numericColumns, textColumns and excludeColumns + var overallColumns = new string[][] { catalogColumns, numericColumns, textColumns, excludeColumns } + .Where(c => c != null) + .SelectMany(c => c); + + if (overallColumns != null) + { + Contracts.Assert(overallColumns.Count() == overallColumns.Distinct().Count(), "detect overlapping among catalogColumns, numericColumns, textColumns and excludedColumns"); + } + var columnInfo = new ColumnInformation(); if (excludeColumns != null) From 75c43be16f5fc830f3ca01cd7112e94d5bca4b09 Mon Sep 17 00:00:00 2001 From: XiaoYun Zhang Date: Wed, 25 May 2022 13:17:01 -0700 Subject: [PATCH 3/7] make public --- src/Microsoft.ML.AutoML/API/AutoCatalog.cs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs index 239bebfbdc..812e7b1890 100644 --- a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs +++ b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs @@ -638,7 +638,7 @@ public MultiModelPipeline Featurizer(IDataView data, string outputColumnName = " /// column information. /// output feature column. /// - internal MultiModelPipeline Featurizer(IDataView data, ColumnInformation columnInformation, string outputColumnName = "Features") + public MultiModelPipeline Featurizer(IDataView data, ColumnInformation columnInformation, string outputColumnName = "Features") { var columnPurposes = PurposeInference.InferPurposes(this._context, data, columnInformation); var textFeatures = columnPurposes.Where(c => c.Purpose == ColumnPurpose.TextFeature); From 95dc0f784c87f845a9c480c5bcb7a4a301fd8074 Mon Sep 17 00:00:00 2001 From: XiaoYun Zhang Date: Wed, 25 May 2022 13:18:34 -0700 Subject: [PATCH 4/7] update --- src/Microsoft.ML.AutoML/API/AutoCatalog.cs | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs index 812e7b1890..4a5efe92b2 100644 --- a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs +++ b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs @@ -626,6 +626,22 @@ public MultiModelPipeline Featurizer(IDataView data, string outputColumnName = " } } + if (numericColumns != null) + { + foreach (var column in numericColumns) + { + columnInfo.NumericColumnNames.Add(column); + } + } + + if (textColumns != null) + { + foreach (var column in textColumns) + { + columnInfo.TextColumnNames.Add(column); + } + } + return this.Featurizer(data, columnInfo, outputColumnName); } From 871a1e0be1741ab11742fd158e2e41b591802f8b Mon Sep 17 00:00:00 2001 From: Xiaoyun Zhang Date: Wed, 15 Jun 2022 18:58:17 -0700 Subject: [PATCH 5/7] Update AutoCatalog.cs --- src/Microsoft.ML.AutoML/API/AutoCatalog.cs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs index 4a5efe92b2..58e4c42c5b 100644 --- a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs +++ b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs @@ -598,14 +598,14 @@ internal SweepableEstimator[] CatalogFeaturizer(string[] outputColumnNames, stri /// columns that won't be included when featurizing, like label public MultiModelPipeline Featurizer(IDataView data, string outputColumnName = "Features", string[] catalogColumns = null, string[] numericColumns = null, string[] textColumns = null, string[] excludeColumns = null) { - // validate if there's overlapping among catalogColumns, numericColumns, textColumns and excludeColumns + // validate if there's overlapping among catalogColumns, numericColumns, textColumns, and excludeColumns var overallColumns = new string[][] { catalogColumns, numericColumns, textColumns, excludeColumns } .Where(c => c != null) .SelectMany(c => c); if (overallColumns != null) { - Contracts.Assert(overallColumns.Count() == overallColumns.Distinct().Count(), "detect overlapping among catalogColumns, numericColumns, textColumns and excludedColumns"); + Contracts.Assert(overallColumns.Count() == overallColumns.Distinct().Count(), "detected overlapping among catalogColumns, numericColumns, textColumns, and excludedColumns"); } var columnInfo = new ColumnInformation(); From 95fb8b6f725671c70861197dca9c529e06852681 Mon Sep 17 00:00:00 2001 From: Xiaoyun Zhang Date: Wed, 15 Jun 2022 18:59:23 -0700 Subject: [PATCH 6/7] Update AutoCatalog.cs --- src/Microsoft.ML.AutoML/API/AutoCatalog.cs | 1 - 1 file changed, 1 deletion(-) diff --git a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs index 58e4c42c5b..74d622ff72 100644 --- a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs +++ b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs @@ -692,7 +692,6 @@ public MultiModelPipeline Featurizer(IDataView data, ColumnInformation columnInf } return pipeline; - } } } From ff9d47e6248b525a52ea2c4d14a21116e12ed207 Mon Sep 17 00:00:00 2001 From: XiaoYun Zhang Date: Thu, 16 Jun 2022 10:21:32 -0700 Subject: [PATCH 7/7] add null check and add return comment --- src/Microsoft.ML.AutoML/API/AutoCatalog.cs | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs index 74d622ff72..1f78d5e0c3 100644 --- a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs +++ b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs @@ -557,6 +557,8 @@ internal SweepableEstimator[] TextFeaturizer(string outputColumnName, string inp /// input column names. internal SweepableEstimator[] NumericFeaturizer(string[] outputColumnNames, string[] inputColumnNames) { + Contracts.CheckValue(inputColumnNames, nameof(inputColumnNames)); + Contracts.CheckValue(outputColumnNames, nameof(outputColumnNames)); Contracts.Check(outputColumnNames.Count() == inputColumnNames.Count() && outputColumnNames.Count() > 0, "outputColumnNames and inputColumnNames must have the same length and greater than 0"); var replaceMissingValueOption = new ReplaceMissingValueOption { @@ -598,14 +600,16 @@ internal SweepableEstimator[] CatalogFeaturizer(string[] outputColumnNames, stri /// columns that won't be included when featurizing, like label public MultiModelPipeline Featurizer(IDataView data, string outputColumnName = "Features", string[] catalogColumns = null, string[] numericColumns = null, string[] textColumns = null, string[] excludeColumns = null) { - // validate if there's overlapping among catalogColumns, numericColumns, textColumns, and excludeColumns + Contracts.CheckValue(data, nameof(data)); + + // validate if there's overlapping among catalogColumns, numericColumns, textColumns and excludeColumns var overallColumns = new string[][] { catalogColumns, numericColumns, textColumns, excludeColumns } .Where(c => c != null) .SelectMany(c => c); if (overallColumns != null) { - Contracts.Assert(overallColumns.Count() == overallColumns.Distinct().Count(), "detected overlapping among catalogColumns, numericColumns, textColumns, and excludedColumns"); + Contracts.Assert(overallColumns.Count() == overallColumns.Distinct().Count(), "detect overlapping among catalogColumns, numericColumns, textColumns and excludedColumns"); } var columnInfo = new ColumnInformation(); @@ -653,9 +657,12 @@ public MultiModelPipeline Featurizer(IDataView data, string outputColumnName = " /// input data. /// column information. /// output feature column. - /// + /// A for featurization. public MultiModelPipeline Featurizer(IDataView data, ColumnInformation columnInformation, string outputColumnName = "Features") { + Contracts.CheckValue(data, nameof(data)); + Contracts.CheckValue(columnInformation, nameof(columnInformation)); + var columnPurposes = PurposeInference.InferPurposes(this._context, data, columnInformation); var textFeatures = columnPurposes.Where(c => c.Purpose == ColumnPurpose.TextFeature); var numericFeatures = columnPurposes.Where(c => c.Purpose == ColumnPurpose.NumericFeature);