From 9bbccafe5f8701e4e53b7e7758e8c352ea99df99 Mon Sep 17 00:00:00 2001 From: Najeeb Kazmi Date: Mon, 3 Feb 2020 17:45:28 -0800 Subject: [PATCH 1/8] Add Seed property to MLContext and use as default for data splits --- .../Data/IHostEnvironment.cs | 5 +++++ .../Environment/ConsoleEnvironment.cs | 19 +------------------ .../Environment/HostEnvironmentBase.cs | 11 +++++++---- .../DataLoadSave/DataOperationsCatalog.cs | 5 ++--- src/Microsoft.ML.Data/MLContext.cs | 1 + .../Utilities/LocalEnvironment.cs | 2 +- 6 files changed, 17 insertions(+), 26 deletions(-) diff --git a/src/Microsoft.ML.Core/Data/IHostEnvironment.cs b/src/Microsoft.ML.Core/Data/IHostEnvironment.cs index 959b1e940b..3521a02a2c 100644 --- a/src/Microsoft.ML.Core/Data/IHostEnvironment.cs +++ b/src/Microsoft.ML.Core/Data/IHostEnvironment.cs @@ -66,6 +66,11 @@ public interface IHostEnvironment : IChannelProvider, IProgressChannelProvider /// The catalog of loadable components () that are available in this host. /// ComponentCatalog ComponentCatalog { get; } + + /// + /// The seed property that, if assigned, makes components requiring randomness behave deterministically. + /// + int? Seed { get; } } [BestFriend] diff --git a/src/Microsoft.ML.Core/Environment/ConsoleEnvironment.cs b/src/Microsoft.ML.Core/Environment/ConsoleEnvironment.cs index 7204484345..27efcc185d 100644 --- a/src/Microsoft.ML.Core/Environment/ConsoleEnvironment.cs +++ b/src/Microsoft.ML.Core/Environment/ConsoleEnvironment.cs @@ -366,24 +366,7 @@ protected override void Dispose(bool disposing) public ConsoleEnvironment(int? seed = null, bool verbose = false, MessageSensitivity sensitivity = MessageSensitivity.All, TextWriter outWriter = null, TextWriter errWriter = null, TextWriter testWriter = null) - : this(RandomUtils.Create(seed), verbose, sensitivity, outWriter, errWriter, testWriter) - { - } - - // REVIEW: do we really care about custom random? If we do, let's make this ctor public. - /// - /// Create an ML.NET environment for local execution, with console feedback. - /// - /// An custom source of randomness to use in the environment. - /// Set to true for fully verbose logging. - /// Allowed message sensitivity. - /// Text writer to print normal messages to. - /// Text writer to print error messages to. - /// Optional TextWriter to write messages if the host is a test environment. - private ConsoleEnvironment(Random rand, bool verbose = false, - MessageSensitivity sensitivity = MessageSensitivity.All, - TextWriter outWriter = null, TextWriter errWriter = null, TextWriter testWriter = null) - : base(rand, verbose, nameof(ConsoleEnvironment)) + : base(seed, verbose, nameof(ConsoleEnvironment)) { Contracts.CheckValueOrNull(outWriter); Contracts.CheckValueOrNull(errWriter); diff --git a/src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs b/src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs index da0e2e711c..3855c65a3e 100644 --- a/src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs +++ b/src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs @@ -330,6 +330,9 @@ public void RemoveListener(Action listenerFunc) // The random number generator for this host. private readonly Random _rand; + + public int? Seed { get; } + // A dictionary mapping the type of message to the Dispatcher that gets the strongly typed dispatch delegate. protected readonly ConcurrentDictionary ListenerDict; @@ -345,14 +348,14 @@ public void RemoveListener(Action listenerFunc) private readonly List> _children; /// - /// The main constructor. + /// The main constructor. /// - protected HostEnvironmentBase(Random rand, bool verbose, + protected HostEnvironmentBase(int? seed, bool verbose, string shortName = null, string parentFullName = null) : base(shortName, parentFullName, verbose) { - Contracts.CheckValueOrNull(rand); - _rand = rand ?? RandomUtils.Create(); + Seed = seed; + _rand = RandomUtils.Create(Seed); ListenerDict = new ConcurrentDictionary(); ProgressTracker = new ProgressReporting.ProgressTracker(this); _cancelLock = new object(); diff --git a/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs b/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs index 4dd77fd822..7042776492 100644 --- a/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs +++ b/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs @@ -496,13 +496,12 @@ internal static IEnumerable CrossValidationSplit(IHostEnvironment internal static void EnsureGroupPreservationColumn(IHostEnvironment env, ref IDataView data, ref string samplingKeyColumn, int? seed = null) { Contracts.CheckValue(env, nameof(env)); - var host = env.Register("rand"); // We need to handle two cases: if samplingKeyColumn is provided, we use hashJoin to // build a single hash of it. If it is not, we generate a random number. if (samplingKeyColumn == null) { samplingKeyColumn = data.Schema.GetTempColumnName("SamplingKeyColumn"); - data = new GenerateNumberTransform(env, data, samplingKeyColumn, (uint?)(seed ?? host.Rand.Next())); + data = new GenerateNumberTransform(env, data, samplingKeyColumn, (uint?)(seed ?? env.Seed)); } else { @@ -518,7 +517,7 @@ internal static void EnsureGroupPreservationColumn(IHostEnvironment env, ref IDa // instead of having two hash transformations. var origStratCol = samplingKeyColumn; samplingKeyColumn = data.Schema.GetTempColumnName(samplingKeyColumn); - var columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30, (uint)(seed ?? host.Rand.Next())); + var columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30, (uint)(seed ?? env.Seed)); data = new HashingEstimator(env, columnOptions).Fit(data).Transform(data); } else diff --git a/src/Microsoft.ML.Data/MLContext.cs b/src/Microsoft.ML.Data/MLContext.cs index 5969552d53..3c533eb6fd 100644 --- a/src/Microsoft.ML.Data/MLContext.cs +++ b/src/Microsoft.ML.Data/MLContext.cs @@ -140,6 +140,7 @@ private void ProcessMessage(IMessageSource source, ChannelMessage message) IChannel IChannelProvider.Start(string name) => _env.Start(name); IPipe IChannelProvider.StartPipe(string name) => _env.StartPipe(name); IProgressChannel IProgressChannelProvider.StartProgressChannel(string name) => _env.StartProgressChannel(name); + int? IHostEnvironment.Seed => _env.Seed; [BestFriend] internal void CancelExecution() => ((ICancelable)_env).CancelExecution(); diff --git a/src/Microsoft.ML.Data/Utilities/LocalEnvironment.cs b/src/Microsoft.ML.Data/Utilities/LocalEnvironment.cs index f2ca816e70..411c2268bd 100644 --- a/src/Microsoft.ML.Data/Utilities/LocalEnvironment.cs +++ b/src/Microsoft.ML.Data/Utilities/LocalEnvironment.cs @@ -47,7 +47,7 @@ protected override void Dispose(bool disposing) /// /// Random seed. Set to null for a non-deterministic environment. public LocalEnvironment(int? seed = null) - : base(RandomUtils.Create(seed), verbose: false) + : base(seed, verbose: false) { } From 9fd29aca95da8f4359c1ef551ce4f20050327da6 Mon Sep 17 00:00:00 2001 From: Najeeb Kazmi Date: Tue, 4 Feb 2020 14:24:00 -0800 Subject: [PATCH 2/8] Separate Seed property out into an internal interface --- src/Microsoft.ML.Core/Data/IHostEnvironment.cs | 14 +++++++++----- .../Environment/HostEnvironmentBase.cs | 2 +- .../DataLoadSave/DataOperationsCatalog.cs | 6 +++--- src/Microsoft.ML.Data/MLContext.cs | 4 ++-- src/Microsoft.ML.Data/TrainCatalog.cs | 18 +++++++++--------- .../PermutationFeatureImportance.cs | 12 ++++++------ .../RecommenderCatalog.cs | 2 +- 7 files changed, 31 insertions(+), 27 deletions(-) diff --git a/src/Microsoft.ML.Core/Data/IHostEnvironment.cs b/src/Microsoft.ML.Core/Data/IHostEnvironment.cs index 3521a02a2c..f59a37bef6 100644 --- a/src/Microsoft.ML.Core/Data/IHostEnvironment.cs +++ b/src/Microsoft.ML.Core/Data/IHostEnvironment.cs @@ -66,11 +66,6 @@ public interface IHostEnvironment : IChannelProvider, IProgressChannelProvider /// The catalog of loadable components () that are available in this host. /// ComponentCatalog ComponentCatalog { get; } - - /// - /// The seed property that, if assigned, makes components requiring randomness behave deterministically. - /// - int? Seed { get; } } [BestFriend] @@ -87,6 +82,15 @@ internal interface ICancelable bool IsCanceled { get; } } + [BestFriend] + internal interface ISeededEnvironment : IHostEnvironment + { + /// + /// The seed property that, if assigned, makes components requiring randomness behave deterministically. + /// + int? Seed { get; } + } + /// /// A host is coupled to a component and provides random number generation and concurrency guidance. /// Note that the random number generation, like the host environment methods, should be accessed only diff --git a/src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs b/src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs index 3855c65a3e..72783d3950 100644 --- a/src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs +++ b/src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs @@ -93,7 +93,7 @@ internal interface IMessageSource /// query progress. /// [BestFriend] - internal abstract class HostEnvironmentBase : ChannelProviderBase, IHostEnvironment, IChannelProvider, ICancelable + internal abstract class HostEnvironmentBase : ChannelProviderBase, ISeededEnvironment, IHostEnvironment, IChannelProvider, ICancelable where TEnv : HostEnvironmentBase { void ICancelable.CancelExecution() diff --git a/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs b/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs index 7042776492..f1455c92d3 100644 --- a/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs +++ b/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs @@ -17,7 +17,7 @@ namespace Microsoft.ML public sealed class DataOperationsCatalog : IInternalCatalog { IHostEnvironment IInternalCatalog.Environment => _env; - private readonly IHostEnvironment _env; + private readonly ISeededEnvironment _env; /// /// A pair of datasets, for the train and test set. @@ -44,7 +44,7 @@ internal TrainTestData(IDataView trainSet, IDataView testSet) } } - internal DataOperationsCatalog(IHostEnvironment env) + internal DataOperationsCatalog(ISeededEnvironment env) { Contracts.AssertValue(env); _env = env; @@ -493,7 +493,7 @@ internal static IEnumerable CrossValidationSplit(IHostEnvironment /// /// Ensures the provided is valid for , hashing it if necessary, or creates a new column is null. /// - internal static void EnsureGroupPreservationColumn(IHostEnvironment env, ref IDataView data, ref string samplingKeyColumn, int? seed = null) + internal static void EnsureGroupPreservationColumn(ISeededEnvironment env, ref IDataView data, ref string samplingKeyColumn, int? seed = null) { Contracts.CheckValue(env, nameof(env)); // We need to handle two cases: if samplingKeyColumn is provided, we use hashJoin to diff --git a/src/Microsoft.ML.Data/MLContext.cs b/src/Microsoft.ML.Data/MLContext.cs index 3c533eb6fd..7131b46715 100644 --- a/src/Microsoft.ML.Data/MLContext.cs +++ b/src/Microsoft.ML.Data/MLContext.cs @@ -14,7 +14,7 @@ namespace Microsoft.ML /// create components for data preparation, feature enginering, training, prediction, model evaluation. /// It also allows logging, execution control, and the ability set repeatable random numbers. /// - public sealed class MLContext : IHostEnvironment + public sealed class MLContext : ISeededEnvironment, IHostEnvironment { // REVIEW: consider making LocalEnvironment and MLContext the same class instead of encapsulation. private readonly LocalEnvironment _env; @@ -140,7 +140,7 @@ private void ProcessMessage(IMessageSource source, ChannelMessage message) IChannel IChannelProvider.Start(string name) => _env.Start(name); IPipe IChannelProvider.StartPipe(string name) => _env.StartPipe(name); IProgressChannel IProgressChannelProvider.StartProgressChannel(string name) => _env.StartProgressChannel(name); - int? IHostEnvironment.Seed => _env.Seed; + int? ISeededEnvironment.Seed => _env.Seed; [BestFriend] internal void CancelExecution() => ((ICancelable)_env).CancelExecution(); diff --git a/src/Microsoft.ML.Data/TrainCatalog.cs b/src/Microsoft.ML.Data/TrainCatalog.cs index cedbb91627..7d16e5749a 100644 --- a/src/Microsoft.ML.Data/TrainCatalog.cs +++ b/src/Microsoft.ML.Data/TrainCatalog.cs @@ -20,7 +20,7 @@ public abstract class TrainCatalogBase : IInternalCatalog IHostEnvironment IInternalCatalog.Environment => Environment; [BestFriend] - private protected IHostEnvironment Environment { get; } + private protected ISeededEnvironment Environment { get; } /// /// Results for specific cross-validation fold. @@ -111,7 +111,7 @@ private protected CrossValidationResult[] CrossValidateTrain(IDataView data, IEs } [BestFriend] - private protected TrainCatalogBase(IHostEnvironment env, string registrationName) + private protected TrainCatalogBase(ISeededEnvironment env, string registrationName) { Contracts.CheckValue(env, nameof(env)); env.CheckNonEmpty(registrationName, nameof(registrationName)); @@ -151,7 +151,7 @@ public sealed class BinaryClassificationCatalog : TrainCatalogBase /// public BinaryClassificationTrainers Trainers { get; } - internal BinaryClassificationCatalog(IHostEnvironment env) + internal BinaryClassificationCatalog(ISeededEnvironment env) : base(env, nameof(BinaryClassificationCatalog)) { Calibrators = new CalibratorsCatalog(this); @@ -388,7 +388,7 @@ public sealed class ClusteringCatalog : TrainCatalogBase /// /// The clustering context. /// - internal ClusteringCatalog(IHostEnvironment env) + internal ClusteringCatalog(ISeededEnvironment env) : base(env, nameof(ClusteringCatalog)) { Trainers = new ClusteringTrainers(this); @@ -468,7 +468,7 @@ public sealed class MulticlassClassificationCatalog : TrainCatalogBase /// public MulticlassClassificationTrainers Trainers { get; } - internal MulticlassClassificationCatalog(IHostEnvironment env) + internal MulticlassClassificationCatalog(ISeededEnvironment env) : base(env, nameof(MulticlassClassificationCatalog)) { Trainers = new MulticlassClassificationTrainers(this); @@ -549,7 +549,7 @@ public sealed class RegressionCatalog : TrainCatalogBase /// public RegressionTrainers Trainers { get; } - internal RegressionCatalog(IHostEnvironment env) + internal RegressionCatalog(ISeededEnvironment env) : base(env, nameof(RegressionCatalog)) { Trainers = new RegressionTrainers(this); @@ -619,7 +619,7 @@ public sealed class RankingCatalog : TrainCatalogBase /// public RankingTrainers Trainers { get; } - internal RankingCatalog(IHostEnvironment env) + internal RankingCatalog(ISeededEnvironment env) : base(env, nameof(RankingCatalog)) { Trainers = new RankingTrainers(this); @@ -685,7 +685,7 @@ public sealed class AnomalyDetectionCatalog : TrainCatalogBase /// public AnomalyDetectionTrainers Trainers { get; } - internal AnomalyDetectionCatalog(IHostEnvironment env) + internal AnomalyDetectionCatalog(ISeededEnvironment env) : base(env, nameof(AnomalyDetectionCatalog)) { Trainers = new AnomalyDetectionTrainers(this); @@ -753,7 +753,7 @@ public sealed class ForecastingCatalog : TrainCatalogBase /// public Forecasters Trainers { get; } - internal ForecastingCatalog(IHostEnvironment env) : base(env, nameof(ForecastingCatalog)) + internal ForecastingCatalog(ISeededEnvironment env) : base(env, nameof(ForecastingCatalog)) { Trainers = new Forecasters(this); } diff --git a/src/Microsoft.ML.EntryPoints/PermutationFeatureImportance.cs b/src/Microsoft.ML.EntryPoints/PermutationFeatureImportance.cs index fc620beeb4..ce98552f33 100644 --- a/src/Microsoft.ML.EntryPoints/PermutationFeatureImportance.cs +++ b/src/Microsoft.ML.EntryPoints/PermutationFeatureImportance.cs @@ -19,7 +19,7 @@ namespace Microsoft.ML.Transforms internal static class PermutationFeatureImportanceEntryPoints { [TlcModule.EntryPoint(Name = "Transforms.PermutationFeatureImportance", Desc = "Permutation Feature Importance (PFI)", UserName = "PFI", ShortName = "PFI")] - public static PermutationFeatureImportanceOutput PermutationFeatureImportance(IHostEnvironment env, PermutationFeatureImportanceArguments input) + public static PermutationFeatureImportanceOutput PermutationFeatureImportance(ISeededEnvironment env, PermutationFeatureImportanceArguments input) { Contracts.CheckValue(env, nameof(env)); var host = env.Register("Pfi"); @@ -57,7 +57,7 @@ internal sealed class PermutationFeatureImportanceArguments : TransformInputBase internal static class PermutationFeatureImportanceUtils { internal static IDataView GetMetrics( - IHostEnvironment env, + ISeededEnvironment env, IPredictor predictor, RoleMappedData roleMappedData, PermutationFeatureImportanceArguments input) @@ -82,7 +82,7 @@ internal static IDataView GetMetrics( } private static IDataView GetBinaryMetrics( - IHostEnvironment env, + ISeededEnvironment env, IPredictor predictor, RoleMappedData roleMappedData, PermutationFeatureImportanceArguments input) @@ -139,7 +139,7 @@ private static IDataView GetBinaryMetrics( } private static IDataView GetMulticlassMetrics( - IHostEnvironment env, + ISeededEnvironment env, IPredictor predictor, RoleMappedData roleMappedData, PermutationFeatureImportanceArguments input) @@ -198,7 +198,7 @@ private static IDataView GetMulticlassMetrics( } private static IDataView GetRegressionMetrics( - IHostEnvironment env, + ISeededEnvironment env, IPredictor predictor, RoleMappedData roleMappedData, PermutationFeatureImportanceArguments input) @@ -249,7 +249,7 @@ private static IDataView GetRegressionMetrics( } private static IDataView GetRankingMetrics( - IHostEnvironment env, + ISeededEnvironment env, IPredictor predictor, RoleMappedData roleMappedData, PermutationFeatureImportanceArguments input) diff --git a/src/Microsoft.ML.Recommender/RecommenderCatalog.cs b/src/Microsoft.ML.Recommender/RecommenderCatalog.cs index 5aa29e1c81..df4ca44325 100644 --- a/src/Microsoft.ML.Recommender/RecommenderCatalog.cs +++ b/src/Microsoft.ML.Recommender/RecommenderCatalog.cs @@ -29,7 +29,7 @@ public sealed class RecommendationCatalog : TrainCatalogBase /// public RecommendationTrainers Trainers { get; } - internal RecommendationCatalog(IHostEnvironment env) + internal RecommendationCatalog(ISeededEnvironment env) : base(env, nameof(RecommendationCatalog)) { Trainers = new RecommendationTrainers(this); From 2080dc0373e170715a450b639eee309e1d2be988 Mon Sep 17 00:00:00 2001 From: Najeeb Kazmi Date: Tue, 4 Feb 2020 15:15:39 -0800 Subject: [PATCH 3/8] Change seed in AutoFitMultiTest --- test/Microsoft.ML.AutoML.Tests/AutoFitTests.cs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/Microsoft.ML.AutoML.Tests/AutoFitTests.cs b/test/Microsoft.ML.AutoML.Tests/AutoFitTests.cs index 1c51bb4c0e..ac4f6f95c7 100644 --- a/test/Microsoft.ML.AutoML.Tests/AutoFitTests.cs +++ b/test/Microsoft.ML.AutoML.Tests/AutoFitTests.cs @@ -37,7 +37,7 @@ public void AutoFitBinaryTest() [Fact] public void AutoFitMultiTest() { - var context = new MLContext(1); + var context = new MLContext(42); var columnInference = context.Auto().InferColumns(DatasetUtil.TrivialMulticlassDatasetPath, DatasetUtil.TrivialMulticlassDatasetLabel); var textLoader = context.Data.CreateTextLoader(columnInference.TextLoaderOptions); var trainData = textLoader.Load(DatasetUtil.TrivialMulticlassDatasetPath); From 218bebd71b5013e826c42adaeb505a6b9fc31c92 Mon Sep 17 00:00:00 2001 From: Najeeb Kazmi Date: Tue, 4 Feb 2020 16:47:14 -0800 Subject: [PATCH 4/8] Check typeof ISeededEnvironment in ComponentCatalog --- src/Microsoft.ML.Core/ComponentModel/ComponentCatalog.cs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/Microsoft.ML.Core/ComponentModel/ComponentCatalog.cs b/src/Microsoft.ML.Core/ComponentModel/ComponentCatalog.cs index 4a7e211007..8557cdd2fa 100644 --- a/src/Microsoft.ML.Core/ComponentModel/ComponentCatalog.cs +++ b/src/Microsoft.ML.Core/ComponentModel/ComponentCatalog.cs @@ -355,7 +355,7 @@ internal EntryPointInfo(MethodInfo method, var parameters = method.GetParameters(); if (parameters.Length != 2 && parameters.Length != 3) throw Contracts.Except("Method '{0}' has {1} parameters, but must have 2 or 3", method.Name, parameters.Length); - if (parameters[0].ParameterType != typeof(IHostEnvironment)) + if (parameters[0].ParameterType != typeof(IHostEnvironment) && parameters[0].ParameterType != typeof(ISeededEnvironment)) throw Contracts.Except("Method '{0}', 1st parameter is {1}, but must be IHostEnvironment", method.Name, parameters[0].ParameterType); InputType = parameters[1].ParameterType; var outputType = method.ReturnType; From d098feeedcb4e459b02476c44ab87f6659376eb8 Mon Sep 17 00:00:00 2001 From: Najeeb Kazmi Date: Wed, 5 Feb 2020 14:30:46 -0800 Subject: [PATCH 5/8] PR feedback --- .../ComponentModel/ComponentCatalog.cs | 2 +- .../DataLoadSave/DataOperationsCatalog.cs | 10 +++++----- src/Microsoft.ML.Data/MLContext.cs | 2 +- src/Microsoft.ML.Data/TrainCatalog.cs | 18 +++++++++--------- .../PermutationFeatureImportance.cs | 12 ++++++------ .../RecommenderCatalog.cs | 2 +- 6 files changed, 23 insertions(+), 23 deletions(-) diff --git a/src/Microsoft.ML.Core/ComponentModel/ComponentCatalog.cs b/src/Microsoft.ML.Core/ComponentModel/ComponentCatalog.cs index 8557cdd2fa..4a7e211007 100644 --- a/src/Microsoft.ML.Core/ComponentModel/ComponentCatalog.cs +++ b/src/Microsoft.ML.Core/ComponentModel/ComponentCatalog.cs @@ -355,7 +355,7 @@ internal EntryPointInfo(MethodInfo method, var parameters = method.GetParameters(); if (parameters.Length != 2 && parameters.Length != 3) throw Contracts.Except("Method '{0}' has {1} parameters, but must have 2 or 3", method.Name, parameters.Length); - if (parameters[0].ParameterType != typeof(IHostEnvironment) && parameters[0].ParameterType != typeof(ISeededEnvironment)) + if (parameters[0].ParameterType != typeof(IHostEnvironment)) throw Contracts.Except("Method '{0}', 1st parameter is {1}, but must be IHostEnvironment", method.Name, parameters[0].ParameterType); InputType = parameters[1].ParameterType; var outputType = method.ReturnType; diff --git a/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs b/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs index f1455c92d3..fcb43a8a3e 100644 --- a/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs +++ b/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs @@ -17,7 +17,7 @@ namespace Microsoft.ML public sealed class DataOperationsCatalog : IInternalCatalog { IHostEnvironment IInternalCatalog.Environment => _env; - private readonly ISeededEnvironment _env; + private readonly IHostEnvironment _env; /// /// A pair of datasets, for the train and test set. @@ -44,7 +44,7 @@ internal TrainTestData(IDataView trainSet, IDataView testSet) } } - internal DataOperationsCatalog(ISeededEnvironment env) + internal DataOperationsCatalog(IHostEnvironment env) { Contracts.AssertValue(env); _env = env; @@ -493,7 +493,7 @@ internal static IEnumerable CrossValidationSplit(IHostEnvironment /// /// Ensures the provided is valid for , hashing it if necessary, or creates a new column is null. /// - internal static void EnsureGroupPreservationColumn(ISeededEnvironment env, ref IDataView data, ref string samplingKeyColumn, int? seed = null) + internal static void EnsureGroupPreservationColumn(IHostEnvironment env, ref IDataView data, ref string samplingKeyColumn, int? seed = null) { Contracts.CheckValue(env, nameof(env)); // We need to handle two cases: if samplingKeyColumn is provided, we use hashJoin to @@ -501,7 +501,7 @@ internal static void EnsureGroupPreservationColumn(ISeededEnvironment env, ref I if (samplingKeyColumn == null) { samplingKeyColumn = data.Schema.GetTempColumnName("SamplingKeyColumn"); - data = new GenerateNumberTransform(env, data, samplingKeyColumn, (uint?)(seed ?? env.Seed)); + data = new GenerateNumberTransform(env, data, samplingKeyColumn, (uint?)(seed ?? ((ISeededEnvironment)env).Seed)); } else { @@ -517,7 +517,7 @@ internal static void EnsureGroupPreservationColumn(ISeededEnvironment env, ref I // instead of having two hash transformations. var origStratCol = samplingKeyColumn; samplingKeyColumn = data.Schema.GetTempColumnName(samplingKeyColumn); - var columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30, (uint)(seed ?? env.Seed)); + var columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30, (uint)(seed ?? ((ISeededEnvironment)env).Seed)); data = new HashingEstimator(env, columnOptions).Fit(data).Transform(data); } else diff --git a/src/Microsoft.ML.Data/MLContext.cs b/src/Microsoft.ML.Data/MLContext.cs index 7131b46715..ccb708addc 100644 --- a/src/Microsoft.ML.Data/MLContext.cs +++ b/src/Microsoft.ML.Data/MLContext.cs @@ -14,7 +14,7 @@ namespace Microsoft.ML /// create components for data preparation, feature enginering, training, prediction, model evaluation. /// It also allows logging, execution control, and the ability set repeatable random numbers. /// - public sealed class MLContext : ISeededEnvironment, IHostEnvironment + public sealed class MLContext : ISeededEnvironment { // REVIEW: consider making LocalEnvironment and MLContext the same class instead of encapsulation. private readonly LocalEnvironment _env; diff --git a/src/Microsoft.ML.Data/TrainCatalog.cs b/src/Microsoft.ML.Data/TrainCatalog.cs index 7d16e5749a..cedbb91627 100644 --- a/src/Microsoft.ML.Data/TrainCatalog.cs +++ b/src/Microsoft.ML.Data/TrainCatalog.cs @@ -20,7 +20,7 @@ public abstract class TrainCatalogBase : IInternalCatalog IHostEnvironment IInternalCatalog.Environment => Environment; [BestFriend] - private protected ISeededEnvironment Environment { get; } + private protected IHostEnvironment Environment { get; } /// /// Results for specific cross-validation fold. @@ -111,7 +111,7 @@ private protected CrossValidationResult[] CrossValidateTrain(IDataView data, IEs } [BestFriend] - private protected TrainCatalogBase(ISeededEnvironment env, string registrationName) + private protected TrainCatalogBase(IHostEnvironment env, string registrationName) { Contracts.CheckValue(env, nameof(env)); env.CheckNonEmpty(registrationName, nameof(registrationName)); @@ -151,7 +151,7 @@ public sealed class BinaryClassificationCatalog : TrainCatalogBase /// public BinaryClassificationTrainers Trainers { get; } - internal BinaryClassificationCatalog(ISeededEnvironment env) + internal BinaryClassificationCatalog(IHostEnvironment env) : base(env, nameof(BinaryClassificationCatalog)) { Calibrators = new CalibratorsCatalog(this); @@ -388,7 +388,7 @@ public sealed class ClusteringCatalog : TrainCatalogBase /// /// The clustering context. /// - internal ClusteringCatalog(ISeededEnvironment env) + internal ClusteringCatalog(IHostEnvironment env) : base(env, nameof(ClusteringCatalog)) { Trainers = new ClusteringTrainers(this); @@ -468,7 +468,7 @@ public sealed class MulticlassClassificationCatalog : TrainCatalogBase /// public MulticlassClassificationTrainers Trainers { get; } - internal MulticlassClassificationCatalog(ISeededEnvironment env) + internal MulticlassClassificationCatalog(IHostEnvironment env) : base(env, nameof(MulticlassClassificationCatalog)) { Trainers = new MulticlassClassificationTrainers(this); @@ -549,7 +549,7 @@ public sealed class RegressionCatalog : TrainCatalogBase /// public RegressionTrainers Trainers { get; } - internal RegressionCatalog(ISeededEnvironment env) + internal RegressionCatalog(IHostEnvironment env) : base(env, nameof(RegressionCatalog)) { Trainers = new RegressionTrainers(this); @@ -619,7 +619,7 @@ public sealed class RankingCatalog : TrainCatalogBase /// public RankingTrainers Trainers { get; } - internal RankingCatalog(ISeededEnvironment env) + internal RankingCatalog(IHostEnvironment env) : base(env, nameof(RankingCatalog)) { Trainers = new RankingTrainers(this); @@ -685,7 +685,7 @@ public sealed class AnomalyDetectionCatalog : TrainCatalogBase /// public AnomalyDetectionTrainers Trainers { get; } - internal AnomalyDetectionCatalog(ISeededEnvironment env) + internal AnomalyDetectionCatalog(IHostEnvironment env) : base(env, nameof(AnomalyDetectionCatalog)) { Trainers = new AnomalyDetectionTrainers(this); @@ -753,7 +753,7 @@ public sealed class ForecastingCatalog : TrainCatalogBase /// public Forecasters Trainers { get; } - internal ForecastingCatalog(ISeededEnvironment env) : base(env, nameof(ForecastingCatalog)) + internal ForecastingCatalog(IHostEnvironment env) : base(env, nameof(ForecastingCatalog)) { Trainers = new Forecasters(this); } diff --git a/src/Microsoft.ML.EntryPoints/PermutationFeatureImportance.cs b/src/Microsoft.ML.EntryPoints/PermutationFeatureImportance.cs index ce98552f33..fc620beeb4 100644 --- a/src/Microsoft.ML.EntryPoints/PermutationFeatureImportance.cs +++ b/src/Microsoft.ML.EntryPoints/PermutationFeatureImportance.cs @@ -19,7 +19,7 @@ namespace Microsoft.ML.Transforms internal static class PermutationFeatureImportanceEntryPoints { [TlcModule.EntryPoint(Name = "Transforms.PermutationFeatureImportance", Desc = "Permutation Feature Importance (PFI)", UserName = "PFI", ShortName = "PFI")] - public static PermutationFeatureImportanceOutput PermutationFeatureImportance(ISeededEnvironment env, PermutationFeatureImportanceArguments input) + public static PermutationFeatureImportanceOutput PermutationFeatureImportance(IHostEnvironment env, PermutationFeatureImportanceArguments input) { Contracts.CheckValue(env, nameof(env)); var host = env.Register("Pfi"); @@ -57,7 +57,7 @@ internal sealed class PermutationFeatureImportanceArguments : TransformInputBase internal static class PermutationFeatureImportanceUtils { internal static IDataView GetMetrics( - ISeededEnvironment env, + IHostEnvironment env, IPredictor predictor, RoleMappedData roleMappedData, PermutationFeatureImportanceArguments input) @@ -82,7 +82,7 @@ internal static IDataView GetMetrics( } private static IDataView GetBinaryMetrics( - ISeededEnvironment env, + IHostEnvironment env, IPredictor predictor, RoleMappedData roleMappedData, PermutationFeatureImportanceArguments input) @@ -139,7 +139,7 @@ private static IDataView GetBinaryMetrics( } private static IDataView GetMulticlassMetrics( - ISeededEnvironment env, + IHostEnvironment env, IPredictor predictor, RoleMappedData roleMappedData, PermutationFeatureImportanceArguments input) @@ -198,7 +198,7 @@ private static IDataView GetMulticlassMetrics( } private static IDataView GetRegressionMetrics( - ISeededEnvironment env, + IHostEnvironment env, IPredictor predictor, RoleMappedData roleMappedData, PermutationFeatureImportanceArguments input) @@ -249,7 +249,7 @@ private static IDataView GetRegressionMetrics( } private static IDataView GetRankingMetrics( - ISeededEnvironment env, + IHostEnvironment env, IPredictor predictor, RoleMappedData roleMappedData, PermutationFeatureImportanceArguments input) diff --git a/src/Microsoft.ML.Recommender/RecommenderCatalog.cs b/src/Microsoft.ML.Recommender/RecommenderCatalog.cs index df4ca44325..5aa29e1c81 100644 --- a/src/Microsoft.ML.Recommender/RecommenderCatalog.cs +++ b/src/Microsoft.ML.Recommender/RecommenderCatalog.cs @@ -29,7 +29,7 @@ public sealed class RecommendationCatalog : TrainCatalogBase /// public RecommendationTrainers Trainers { get; } - internal RecommendationCatalog(ISeededEnvironment env) + internal RecommendationCatalog(IHostEnvironment env) : base(env, nameof(RecommendationCatalog)) { Trainers = new RecommendationTrainers(this); From 66fc8491c669983ad8ef2ec96fdee0eede83131a Mon Sep 17 00:00:00 2001 From: Najeeb Kazmi Date: Wed, 5 Feb 2020 14:32:33 -0800 Subject: [PATCH 6/8] nit --- src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs b/src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs index 72783d3950..f776b08f56 100644 --- a/src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs +++ b/src/Microsoft.ML.Core/Environment/HostEnvironmentBase.cs @@ -93,7 +93,7 @@ internal interface IMessageSource /// query progress. /// [BestFriend] - internal abstract class HostEnvironmentBase : ChannelProviderBase, ISeededEnvironment, IHostEnvironment, IChannelProvider, ICancelable + internal abstract class HostEnvironmentBase : ChannelProviderBase, ISeededEnvironment, IChannelProvider, ICancelable where TEnv : HostEnvironmentBase { void ICancelable.CancelExecution() From 0ee3b9ab1a97bfd9e0c63f04089fe1b22ecb98cc Mon Sep 17 00:00:00 2001 From: Najeeb Kazmi Date: Thu, 6 Feb 2020 10:21:19 -0800 Subject: [PATCH 7/8] PR feedback --- src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs b/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs index fcb43a8a3e..bba8ce8a1b 100644 --- a/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs +++ b/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs @@ -517,7 +517,7 @@ internal static void EnsureGroupPreservationColumn(IHostEnvironment env, ref IDa // instead of having two hash transformations. var origStratCol = samplingKeyColumn; samplingKeyColumn = data.Schema.GetTempColumnName(samplingKeyColumn); - var columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30, (uint)(seed ?? ((ISeededEnvironment)env).Seed)); + var columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30, (uint?)(seed ?? ((ISeededEnvironment)env).Seed)); data = new HashingEstimator(env, columnOptions).Fit(data).Transform(data); } else From 67d69c05bf0d723d339c9b7a22baadb5bc3c8f87 Mon Sep 17 00:00:00 2001 From: Najeeb Kazmi Date: Thu, 6 Feb 2020 12:55:21 -0800 Subject: [PATCH 8/8] Handle casting of nullable int to uint --- .../DataLoadSave/DataOperationsCatalog.cs | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs b/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs index bba8ce8a1b..4d5246586a 100644 --- a/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs +++ b/src/Microsoft.ML.Data/DataLoadSave/DataOperationsCatalog.cs @@ -517,7 +517,13 @@ internal static void EnsureGroupPreservationColumn(IHostEnvironment env, ref IDa // instead of having two hash transformations. var origStratCol = samplingKeyColumn; samplingKeyColumn = data.Schema.GetTempColumnName(samplingKeyColumn); - var columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30, (uint?)(seed ?? ((ISeededEnvironment)env).Seed)); + HashingEstimator.ColumnOptionsInternal columnOptions; + if (seed.HasValue) + columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30, (uint)seed.Value); + else if (((ISeededEnvironment)env).Seed.HasValue) + columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30, (uint)((ISeededEnvironment)env).Seed.Value); + else + columnOptions = new HashingEstimator.ColumnOptionsInternal(samplingKeyColumn, origStratCol, 30); data = new HashingEstimator(env, columnOptions).Fit(data).Transform(data); } else