diff --git a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs index 8c4642ba70..859247e6c4 100644 --- a/src/Microsoft.ML.AutoML/API/AutoCatalog.cs +++ b/src/Microsoft.ML.AutoML/API/AutoCatalog.cs @@ -7,6 +7,7 @@ using Microsoft.ML.AutoML.CodeGen; using Microsoft.ML.Data; using Microsoft.ML.SearchSpace; +using Microsoft.ML.Trainers.FastTree; namespace Microsoft.ML.AutoML { @@ -286,18 +287,43 @@ public ColumnInferenceResults InferColumns(string path, uint labelColumnIndex, b /// /// Create a sweepable estimator with a custom factory and search space. /// - internal SweepableEstimator CreateSweepableEstimator(Func> factory, SearchSpace ss = null) + public SweepableEstimator CreateSweepableEstimator(Func> factory, SearchSpace ss = null) where T : class, new() { return new SweepableEstimator((MLContext context, Parameter param) => factory(context, param.AsType()), ss); } - internal AutoMLExperiment CreateExperiment() + /// + /// Create an . + /// + public AutoMLExperiment CreateExperiment() { return new AutoMLExperiment(_context, new AutoMLExperiment.AutoMLExperimentSettings()); } - internal SweepableEstimator[] BinaryClassification(string labelColumnName = DefaultColumnNames.Label, string featureColumnName = DefaultColumnNames.Features, string exampleWeightColumnName = null, bool useFastForest = true, bool useLgbm = true, bool useFastTree = true, bool useLbfgs = true, bool useSdca = true, + /// + /// Create a list of for binary classification. + /// + /// label column name. + /// feature column name. + /// example weight column name. + /// true if use fast forest as available trainer. + /// true if use lgbm as available trainer. + /// true if use fast tree as available trainer. + /// true if use lbfgs as available trainer. + /// true if use sdca as available trainer. + /// if provided, use it as initial option for fast tree, otherwise the default option will be used. + /// if provided, use it as initial option for lgbm, otherwise the default option will be used. + /// if provided, use it as initial option for fast forest, otherwise the default option will be used. + /// if provided, use it as initial option for lbfgs, otherwise the default option will be used. + /// if provided, use it as initial option for sdca, otherwise the default option will be used. + /// if provided, use it as search space for fast tree, otherwise the default search space will be used. + /// if provided, use it as search space for lgbm, otherwise the default search space will be used. + /// if provided, use it as search space for fast forest, otherwise the default search space will be used. + /// if provided, use it as search space for lbfgs, otherwise the default search space will be used. + /// if provided, use it as search space for sdca, otherwise the default search space will be used. + /// + public SweepableEstimator[] BinaryClassification(string labelColumnName = DefaultColumnNames.Label, string featureColumnName = DefaultColumnNames.Features, string exampleWeightColumnName = null, bool useFastForest = true, bool useLgbm = true, bool useFastTree = true, bool useLbfgs = true, bool useSdca = true, FastTreeOption fastTreeOption = null, LgbmOption lgbmOption = null, FastForestOption fastForestOption = null, LbfgsOption lbfgsOption = null, SdcaOption sdcaOption = null, SearchSpace fastTreeSearchSpace = null, SearchSpace lgbmSearchSpace = null, SearchSpace fastForestSearchSpace = null, SearchSpace lbfgsSearchSpace = null, SearchSpace sdcaSearchSpace = null) { @@ -351,7 +377,29 @@ internal SweepableEstimator[] BinaryClassification(string labelColumnName = Defa return res.ToArray(); } - internal SweepableEstimator[] MultiClassification(string labelColumnName = DefaultColumnNames.Label, string featureColumnName = DefaultColumnNames.Features, string exampleWeightColumnName = null, bool useFastForest = true, bool useLgbm = true, bool useFastTree = true, bool useLbfgs = true, bool useSdca = true, + /// + /// Create a list of for multiclass classification. + /// + /// label column name. + /// feature column name. + /// example weight column name. + /// true if use fast forest as available trainer. + /// true if use lgbm as available trainer. + /// true if use fast tree as available trainer. + /// true if use lbfgs as available trainer. + /// true if use sdca as available trainer. + /// if provided, use it as initial option for fast tree, otherwise the default option will be used. + /// if provided, use it as initial option for lgbm, otherwise the default option will be used. + /// if provided, use it as initial option for fast forest, otherwise the default option will be used. + /// if provided, use it as initial option for lbfgs, otherwise the default option will be used. + /// if provided, use it as initial option for sdca, otherwise the default option will be used. + /// if provided, use it as search space for fast tree, otherwise the default search space will be used. + /// if provided, use it as search space for lgbm, otherwise the default search space will be used. + /// if provided, use it as search space for fast forest, otherwise the default search space will be used. + /// if provided, use it as search space for lbfgs, otherwise the default search space will be used. + /// if provided, use it as search space for sdca, otherwise the default search space will be used. + /// + public SweepableEstimator[] MultiClassification(string labelColumnName = DefaultColumnNames.Label, string featureColumnName = DefaultColumnNames.Features, string exampleWeightColumnName = null, bool useFastForest = true, bool useLgbm = true, bool useFastTree = true, bool useLbfgs = true, bool useSdca = true, FastTreeOption fastTreeOption = null, LgbmOption lgbmOption = null, FastForestOption fastForestOption = null, LbfgsOption lbfgsOption = null, SdcaOption sdcaOption = null, SearchSpace fastTreeSearchSpace = null, SearchSpace lgbmSearchSpace = null, SearchSpace fastForestSearchSpace = null, SearchSpace lbfgsSearchSpace = null, SearchSpace sdcaSearchSpace = null) { @@ -407,7 +455,29 @@ internal SweepableEstimator[] MultiClassification(string labelColumnName = Defau return res.ToArray(); } - internal SweepableEstimator[] Regression(string labelColumnName = DefaultColumnNames.Label, string featureColumnName = DefaultColumnNames.Features, string exampleWeightColumnName = null, bool useFastForest = true, bool useLgbm = true, bool useFastTree = true, bool useLbfgs = true, bool useSdca = true, + /// + /// Create a list of for regression. + /// + /// label column name. + /// feature column name. + /// example weight column name. + /// true if use fast forest as available trainer. + /// true if use lgbm as available trainer. + /// true if use fast tree as available trainer. + /// true if use lbfgs as available trainer. + /// true if use sdca as available trainer. + /// if provided, use it as initial option for fast tree, otherwise the default option will be used. + /// if provided, use it as initial option for lgbm, otherwise the default option will be used. + /// if provided, use it as initial option for fast forest, otherwise the default option will be used. + /// if provided, use it as initial option for lbfgs, otherwise the default option will be used. + /// if provided, use it as initial option for sdca, otherwise the default option will be used. + /// if provided, use it as search space for fast tree, otherwise the default search space will be used. + /// if provided, use it as search space for lgbm, otherwise the default search space will be used. + /// if provided, use it as search space for fast forest, otherwise the default search space will be used. + /// if provided, use it as search space for lbfgs, otherwise the default search space will be used. + /// if provided, use it as search space for sdca, otherwise the default search space will be used. + /// + public SweepableEstimator[] Regression(string labelColumnName = DefaultColumnNames.Label, string featureColumnName = DefaultColumnNames.Features, string exampleWeightColumnName = null, bool useFastForest = true, bool useLgbm = true, bool useFastTree = true, bool useLbfgs = true, bool useSdca = true, FastTreeOption fastTreeOption = null, LgbmOption lgbmOption = null, FastForestOption fastForestOption = null, LbfgsOption lbfgsOption = null, SdcaOption sdcaOption = null, SearchSpace fastTreeSearchSpace = null, SearchSpace lgbmSearchSpace = null, SearchSpace fastForestSearchSpace = null, SearchSpace lbfgsSearchSpace = null, SearchSpace sdcaSearchSpace = null) { diff --git a/src/Microsoft.ML.AutoML/API/SweepableExtension.cs b/src/Microsoft.ML.AutoML/API/SweepableExtension.cs index 6922a9597d..006e35ee6e 100644 --- a/src/Microsoft.ML.AutoML/API/SweepableExtension.cs +++ b/src/Microsoft.ML.AutoML/API/SweepableExtension.cs @@ -4,7 +4,7 @@ namespace Microsoft.ML.AutoML { - internal static class SweepableExtension + public static class SweepableExtension { public static SweepableEstimatorPipeline Append(this IEstimator estimator, SweepableEstimator estimator1) { diff --git a/src/Microsoft.ML.AutoML/AutoMLExperiment/AutoMLExperiment.cs b/src/Microsoft.ML.AutoML/AutoMLExperiment/AutoMLExperiment.cs index dfa34dcf33..2120c66572 100644 --- a/src/Microsoft.ML.AutoML/AutoMLExperiment/AutoMLExperiment.cs +++ b/src/Microsoft.ML.AutoML/AutoMLExperiment/AutoMLExperiment.cs @@ -12,7 +12,7 @@ namespace Microsoft.ML.AutoML { - internal class AutoMLExperiment + public class AutoMLExperiment { private readonly AutoMLExperimentSettings _settings; private readonly MLContext _context; @@ -52,12 +52,14 @@ public AutoMLExperiment SetTrainingTimeInSeconds(uint trainingTimeInSeconds) public AutoMLExperiment SetDataset(IDataView train, IDataView test) { - _settings.DatasetSettings = new TrainTestDatasetSettings() + var datasetManager = new TrainTestDatasetManager() { TrainDataset = train, TestDataset = test }; + _serviceCollection.AddSingleton(datasetManager); + return this; } @@ -70,12 +72,14 @@ public AutoMLExperiment SetDataset(TrainTestData trainTestSplit) public AutoMLExperiment SetDataset(IDataView dataset, int fold = 10) { - _settings.DatasetSettings = new CrossValidateDatasetSettings() + var datasetManager = new CrossValidateDatasetManager() { Dataset = dataset, Fold = fold, }; + _serviceCollection.AddSingleton(datasetManager); + return this; } @@ -116,8 +120,15 @@ public AutoMLExperiment SetPipeline(MultiModelPipeline pipeline) return this; } - public AutoMLExperiment SetTrialRunnerFactory(ITrialRunnerFactory factory) + public AutoMLExperiment SetIsMaximizeMetric(bool isMaximize) { + _settings.IsMaximizeMetric = isMaximize; + return this; + } + + public AutoMLExperiment SetTrialRunner(ITrialRunner runner) + { + var factory = new CustomRunnerFactory(runner); var descriptor = new ServiceDescriptor(typeof(ITrialRunnerFactory), factory); if (_serviceCollection.Contains(descriptor)) { @@ -146,36 +157,42 @@ public AutoMLExperiment SetPipeline(SweepableEstimatorPipeline pipeline) public AutoMLExperiment SetEvaluateMetric(BinaryClassificationMetric metric, string labelColumn = "label", string predictedColumn = "Predicted") { - _settings.EvaluateMetric = new BinaryMetricSettings() + var metricManager = new BinaryMetricManager() { Metric = metric, PredictedColumn = predictedColumn, LabelColumn = labelColumn, }; + _serviceCollection.AddSingleton(metricManager); + SetIsMaximizeMetric(metricManager.IsMaximize); return this; } public AutoMLExperiment SetEvaluateMetric(MulticlassClassificationMetric metric, string labelColumn = "label", string predictedColumn = "Predicted") { - _settings.EvaluateMetric = new MultiClassMetricSettings() + var metricManager = new MultiClassMetricManager() { Metric = metric, PredictedColumn = predictedColumn, LabelColumn = labelColumn, }; + _serviceCollection.AddSingleton(metricManager); + SetIsMaximizeMetric(metricManager.IsMaximize); return this; } public AutoMLExperiment SetEvaluateMetric(RegressionMetric metric, string labelColumn = "label", string scoreColumn = "Score") { - _settings.EvaluateMetric = new RegressionMetricSettings() + var metricManager = new RegressionMetricManager() { Metric = metric, ScoreColumn = scoreColumn, LabelColumn = labelColumn, }; + _serviceCollection.AddSingleton(metricManager); + SetIsMaximizeMetric(metricManager.IsMaximize); return this; } @@ -224,13 +241,13 @@ private async Task RunAsync(CancellationToken ct) setting = pipelineProposer.Propose(setting); setting = hyperParameterProposer.Propose(setting); monitor.ReportRunningTrial(setting); - var runner = runnerFactory.CreateTrialRunner(setting); - var trialResult = runner.Run(setting); + var runner = runnerFactory.CreateTrialRunner(); + var trialResult = runner.Run(setting, serviceProvider); monitor.ReportCompletedTrial(trialResult); hyperParameterProposer.Update(setting, trialResult); pipelineProposer.Update(setting, trialResult); - var error = _settings.EvaluateMetric.IsMaximize ? 1 - trialResult.Metric : trialResult.Metric; + var error = _settings.IsMaximizeMetric ? 1 - trialResult.Metric : trialResult.Metric; if (error < _bestError) { _bestTrialResult = trialResult; @@ -264,20 +281,16 @@ private async Task RunAsync(CancellationToken ct) private void ValidateSettings() { Contracts.Assert(_settings.MaxExperimentTimeInSeconds > 0, $"{nameof(ExperimentSettings.MaxExperimentTimeInSeconds)} must be larger than 0"); - Contracts.Assert(_settings.DatasetSettings != null, $"{nameof(_settings.DatasetSettings)} must be not null"); - Contracts.Assert(_settings.EvaluateMetric != null, $"{nameof(_settings.EvaluateMetric)} must be not null"); } public class AutoMLExperimentSettings : ExperimentSettings { - public IDatasetSettings DatasetSettings { get; set; } - - public IMetricSettings EvaluateMetric { get; set; } - public MultiModelPipeline Pipeline { get; set; } public int? Seed { get; set; } + + public bool IsMaximizeMetric { get; set; } } } } diff --git a/src/Microsoft.ML.AutoML/AutoMLExperiment/IDatasetSettings.cs b/src/Microsoft.ML.AutoML/AutoMLExperiment/IDatasetManager.cs similarity index 51% rename from src/Microsoft.ML.AutoML/AutoMLExperiment/IDatasetSettings.cs rename to src/Microsoft.ML.AutoML/AutoMLExperiment/IDatasetManager.cs index c551139928..ac057be17d 100644 --- a/src/Microsoft.ML.AutoML/AutoMLExperiment/IDatasetSettings.cs +++ b/src/Microsoft.ML.AutoML/AutoMLExperiment/IDatasetManager.cs @@ -4,18 +4,22 @@ namespace Microsoft.ML.AutoML { - internal interface IDatasetSettings + /// + /// Interface for dataset manager. This interface doesn't include any method or property definition and is used by and other components to retrieve the instance of the actual + /// dataset manager from containers. + /// + public interface IDatasetManager { } - internal class TrainTestDatasetSettings : IDatasetSettings + public class TrainTestDatasetManager : IDatasetManager { public IDataView TrainDataset { get; set; } public IDataView TestDataset { get; set; } } - internal class CrossValidateDatasetSettings : IDatasetSettings + public class CrossValidateDatasetManager : IDatasetManager { public IDataView Dataset { get; set; } diff --git a/src/Microsoft.ML.AutoML/AutoMLExperiment/IMetricSettings.cs b/src/Microsoft.ML.AutoML/AutoMLExperiment/IMetricManager.cs similarity index 88% rename from src/Microsoft.ML.AutoML/AutoMLExperiment/IMetricSettings.cs rename to src/Microsoft.ML.AutoML/AutoMLExperiment/IMetricManager.cs index 2374ba75f4..cee384f7b9 100644 --- a/src/Microsoft.ML.AutoML/AutoMLExperiment/IMetricSettings.cs +++ b/src/Microsoft.ML.AutoML/AutoMLExperiment/IMetricManager.cs @@ -6,12 +6,15 @@ namespace Microsoft.ML.AutoML { - internal interface IMetricSettings + /// + /// Interface for metric manager. + /// + internal interface IMetricManager { bool IsMaximize { get; } } - internal class BinaryMetricSettings : IMetricSettings + internal class BinaryMetricManager : IMetricManager { public BinaryClassificationMetric Metric { get; set; } @@ -33,7 +36,7 @@ internal class BinaryMetricSettings : IMetricSettings }; } - internal class MultiClassMetricSettings : IMetricSettings + internal class MultiClassMetricManager : IMetricManager { public MulticlassClassificationMetric Metric { get; set; } @@ -52,7 +55,7 @@ internal class MultiClassMetricSettings : IMetricSettings }; } - internal class RegressionMetricSettings : IMetricSettings + internal class RegressionMetricManager : IMetricManager { public RegressionMetric Metric { get; set; } diff --git a/src/Microsoft.ML.AutoML/AutoMLExperiment/IMonitor.cs b/src/Microsoft.ML.AutoML/AutoMLExperiment/IMonitor.cs index 64a10e1a4a..1d17908370 100644 --- a/src/Microsoft.ML.AutoML/AutoMLExperiment/IMonitor.cs +++ b/src/Microsoft.ML.AutoML/AutoMLExperiment/IMonitor.cs @@ -8,7 +8,10 @@ namespace Microsoft.ML.AutoML { - internal interface IMonitor + /// + /// instance for monitor, which is used by to report training progress. + /// + public interface IMonitor { void ReportCompletedTrial(TrialResult result); diff --git a/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialResult.cs b/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialResult.cs index fa60ec5d9c..bd3f19d47b 100644 --- a/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialResult.cs +++ b/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialResult.cs @@ -4,7 +4,7 @@ namespace Microsoft.ML.AutoML { - internal class TrialResult + public class TrialResult { public TrialSettings TrialSettings { get; set; } @@ -12,6 +12,8 @@ internal class TrialResult public double Metric { get; set; } + public bool IsMaximize { get; set; } + public double DurationInMilliseconds { get; set; } } } diff --git a/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialRunner.cs b/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialRunner.cs index 899c6ad3c9..3bac432dc7 100644 --- a/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialRunner.cs +++ b/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialRunner.cs @@ -7,24 +7,32 @@ namespace Microsoft.ML.AutoML { - internal interface ITrialRunner + /// + /// interface for all trial runners. + /// + public interface ITrialRunner { - TrialResult Run(TrialSettings settings); + TrialResult Run(TrialSettings settings, IServiceProvider provider = null); } internal class BinaryClassificationCVRunner : ITrialRunner { private readonly MLContext _context; - public BinaryClassificationCVRunner(MLContext context) + private readonly IDatasetManager _datasetManager; + private readonly IMetricManager _metricManager; + + public BinaryClassificationCVRunner(MLContext context, IDatasetManager datasetManager, IMetricManager metricManager) { _context = context; + _datasetManager = datasetManager; + _metricManager = metricManager; } - public TrialResult Run(TrialSettings settings) + public TrialResult Run(TrialSettings settings, IServiceProvider provider) { var rnd = new Random(settings.ExperimentSettings.Seed ?? 0); - if (settings.ExperimentSettings.DatasetSettings is CrossValidateDatasetSettings datasetSettings - && settings.ExperimentSettings.EvaluateMetric is BinaryMetricSettings metricSettings) + if (_datasetManager is CrossValidateDatasetManager datasetSettings + && _metricManager is BinaryMetricManager metricSettings) { var stopWatch = new Stopwatch(); stopWatch.Start(); @@ -54,6 +62,7 @@ public TrialResult Run(TrialSettings settings) Model = model, TrialSettings = settings, DurationInMilliseconds = stopWatch.ElapsedMilliseconds, + IsMaximize = _metricManager.IsMaximize, }; } @@ -64,16 +73,20 @@ public TrialResult Run(TrialSettings settings) internal class BinaryClassificationTrainTestRunner : ITrialRunner { private readonly MLContext _context; - public BinaryClassificationTrainTestRunner(MLContext context) + private readonly IDatasetManager _datasetManager; + private readonly IMetricManager _metricManager; + + public BinaryClassificationTrainTestRunner(MLContext context, IDatasetManager datasetManager, IMetricManager metricManager) { _context = context; + _metricManager = metricManager; + _datasetManager = datasetManager; } - public TrialResult Run(TrialSettings settings) + public TrialResult Run(TrialSettings settings, IServiceProvider provider) { - var rnd = new Random(settings.ExperimentSettings.Seed ?? 0); - if (settings.ExperimentSettings.DatasetSettings is TrainTestDatasetSettings datasetSettings - && settings.ExperimentSettings.EvaluateMetric is BinaryMetricSettings metricSettings) + if (_datasetManager is TrainTestDatasetManager datasetSettings + && _metricManager is BinaryMetricManager metricSettings) { var stopWatch = new Stopwatch(); stopWatch.Start(); @@ -102,6 +115,7 @@ public TrialResult Run(TrialSettings settings) Model = model, TrialSettings = settings, DurationInMilliseconds = stopWatch.ElapsedMilliseconds, + IsMaximize = _metricManager.IsMaximize, }; } @@ -112,15 +126,20 @@ public TrialResult Run(TrialSettings settings) internal class MultiClassificationTrainTestRunner : ITrialRunner { private readonly MLContext _context; - public MultiClassificationTrainTestRunner(MLContext context) + private readonly IDatasetManager _datasetManager; + private readonly IMetricManager _metricManager; + + public MultiClassificationTrainTestRunner(MLContext context, IDatasetManager datasetManager, IMetricManager metricManager) { _context = context; + _metricManager = metricManager; + _datasetManager = datasetManager; } - public TrialResult Run(TrialSettings settings) + public TrialResult Run(TrialSettings settings, IServiceProvider provider) { - if (settings.ExperimentSettings.DatasetSettings is TrainTestDatasetSettings datasetSettings - && settings.ExperimentSettings.EvaluateMetric is MultiClassMetricSettings metricSettings) + if (_datasetManager is TrainTestDatasetManager datasetSettings + && _metricManager is MultiClassMetricManager metricSettings) { var stopWatch = new Stopwatch(); stopWatch.Start(); @@ -149,6 +168,7 @@ public TrialResult Run(TrialSettings settings) Model = model, TrialSettings = settings, DurationInMilliseconds = stopWatch.ElapsedMilliseconds, + IsMaximize = _metricManager.IsMaximize, }; } @@ -159,16 +179,21 @@ public TrialResult Run(TrialSettings settings) internal class MultiClassificationCVRunner : ITrialRunner { private readonly MLContext _context; - public MultiClassificationCVRunner(MLContext context) + private readonly IDatasetManager _datasetManager; + private readonly IMetricManager _metricManager; + + public MultiClassificationCVRunner(MLContext context, IDatasetManager datasetManager, IMetricManager metricManager) { _context = context; + _metricManager = metricManager; + _datasetManager = datasetManager; } - public TrialResult Run(TrialSettings settings) + public TrialResult Run(TrialSettings settings, IServiceProvider provider) { var rnd = new Random(settings.ExperimentSettings.Seed ?? 0); - if (settings.ExperimentSettings.DatasetSettings is CrossValidateDatasetSettings datasetSettings - && settings.ExperimentSettings.EvaluateMetric is MultiClassMetricSettings metricSettings) + if (_datasetManager is CrossValidateDatasetManager datasetSettings + && _metricManager is MultiClassMetricManager metricSettings) { var stopWatch = new Stopwatch(); stopWatch.Start(); @@ -197,6 +222,7 @@ public TrialResult Run(TrialSettings settings) Model = model, TrialSettings = settings, DurationInMilliseconds = stopWatch.ElapsedMilliseconds, + IsMaximize = _metricManager.IsMaximize, }; } @@ -207,15 +233,20 @@ public TrialResult Run(TrialSettings settings) internal class RegressionTrainTestRunner : ITrialRunner { private readonly MLContext _context; - public RegressionTrainTestRunner(MLContext context) + private readonly IDatasetManager _datasetManager; + private readonly IMetricManager _metricManager; + + public RegressionTrainTestRunner(MLContext context, IDatasetManager datasetManager, IMetricManager metricManager) { _context = context; + _metricManager = metricManager; + _datasetManager = datasetManager; } - public TrialResult Run(TrialSettings settings) + public TrialResult Run(TrialSettings settings, IServiceProvider provider) { - if (settings.ExperimentSettings.DatasetSettings is TrainTestDatasetSettings datasetSettings - && settings.ExperimentSettings.EvaluateMetric is RegressionMetricSettings metricSettings) + if (_datasetManager is TrainTestDatasetManager datasetSettings + && _metricManager is RegressionMetricManager metricSettings) { var stopWatch = new Stopwatch(); stopWatch.Start(); @@ -243,6 +274,7 @@ public TrialResult Run(TrialSettings settings) Model = model, TrialSettings = settings, DurationInMilliseconds = stopWatch.ElapsedMilliseconds, + IsMaximize = _metricManager.IsMaximize, }; } @@ -253,16 +285,21 @@ public TrialResult Run(TrialSettings settings) internal class RegressionCVRunner : ITrialRunner { private readonly MLContext _context; - public RegressionCVRunner(MLContext context) + private readonly IDatasetManager _datasetManager; + private readonly IMetricManager _metricManager; + + public RegressionCVRunner(MLContext context, IDatasetManager datasetManager, IMetricManager metricManager) { _context = context; + _metricManager = metricManager; + _datasetManager = datasetManager; } - public TrialResult Run(TrialSettings settings) + public TrialResult Run(TrialSettings settings, IServiceProvider provider) { var rnd = new Random(settings.ExperimentSettings.Seed ?? 0); - if (settings.ExperimentSettings.DatasetSettings is CrossValidateDatasetSettings datasetSettings - && settings.ExperimentSettings.EvaluateMetric is RegressionMetricSettings metricSettings) + if (_datasetManager is CrossValidateDatasetManager datasetSettings + && _metricManager is RegressionMetricManager metricSettings) { var stopWatch = new Stopwatch(); stopWatch.Start(); @@ -290,6 +327,7 @@ public TrialResult Run(TrialSettings settings) Model = model, TrialSettings = settings, DurationInMilliseconds = stopWatch.ElapsedMilliseconds, + IsMaximize = _metricManager.IsMaximize, }; } diff --git a/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialRunnerFactory.cs b/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialRunnerFactory.cs index 18a3dfc4d1..8bba47d321 100644 --- a/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialRunnerFactory.cs +++ b/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialRunnerFactory.cs @@ -8,9 +8,27 @@ #nullable enable namespace Microsoft.ML.AutoML { - internal interface ITrialRunnerFactory + /// + /// interface for trial runner factory. + /// + public interface ITrialRunnerFactory { - ITrialRunner? CreateTrialRunner(TrialSettings settings); + ITrialRunner? CreateTrialRunner(); + } + + internal class CustomRunnerFactory : ITrialRunnerFactory + { + private readonly ITrialRunner _instance; + + public CustomRunnerFactory(ITrialRunner runner) + { + _instance = runner; + } + + public ITrialRunner? CreateTrialRunner() + { + return _instance; + } } internal class TrialRunnerFactory : ITrialRunnerFactory @@ -22,16 +40,19 @@ public TrialRunnerFactory(IServiceProvider provider) _provider = provider; } - public ITrialRunner? CreateTrialRunner(TrialSettings settings) + public ITrialRunner? CreateTrialRunner() { - ITrialRunner? runner = (settings.ExperimentSettings.DatasetSettings, settings.ExperimentSettings.EvaluateMetric) switch + var datasetManager = _provider.GetService(); + var metricManager = _provider.GetService(); + + ITrialRunner? runner = (datasetManager, metricManager) switch { - (CrossValidateDatasetSettings, BinaryMetricSettings) => _provider.GetService(), - (TrainTestDatasetSettings, BinaryMetricSettings) => _provider.GetService(), - (CrossValidateDatasetSettings, MultiClassMetricSettings) => _provider.GetService(), - (TrainTestDatasetSettings, MultiClassMetricSettings) => _provider.GetService(), - (CrossValidateDatasetSettings, RegressionMetricSettings) => _provider.GetService(), - (TrainTestDatasetSettings, RegressionMetricSettings) => _provider.GetService(), + (CrossValidateDatasetManager, BinaryMetricManager) => _provider.GetService(), + (TrainTestDatasetManager, BinaryMetricManager) => _provider.GetService(), + (CrossValidateDatasetManager, MultiClassMetricManager) => _provider.GetService(), + (TrainTestDatasetManager, MultiClassMetricManager) => _provider.GetService(), + (CrossValidateDatasetManager, RegressionMetricManager) => _provider.GetService(), + (TrainTestDatasetManager, RegressionMetricManager) => _provider.GetService(), _ => throw new NotImplementedException(), }; diff --git a/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialSettings.cs b/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialSettings.cs index 8cd8a2e2ab..19294ffde9 100644 --- a/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialSettings.cs +++ b/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialSettings.cs @@ -6,7 +6,7 @@ namespace Microsoft.ML.AutoML { - internal class TrialSettings + public class TrialSettings { public int TrialId { get; set; } diff --git a/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialSettingsProposer/PipelineProposer.cs b/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialSettingsProposer/PipelineProposer.cs index b157117054..72c175dbdd 100644 --- a/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialSettingsProposer/PipelineProposer.cs +++ b/src/Microsoft.ML.AutoML/AutoMLExperiment/TrialSettingsProposer/PipelineProposer.cs @@ -157,7 +157,7 @@ public void LoadStatusFromFile(string fileName) public void Update(TrialSettings parameter, TrialResult result) { var schema = parameter.Schema; - var error = CaculateError(result.Metric, result.TrialSettings.ExperimentSettings.EvaluateMetric.IsMaximize); + var error = CaculateError(result.Metric, result.IsMaximize); var duration = result.DurationInMilliseconds / 1000; var pipelineIds = _multiModelPipeline.PipelineIds; var isSuccess = duration != 0; diff --git a/src/Microsoft.ML.AutoML/AutoMLExperiment/TunerFactory.cs b/src/Microsoft.ML.AutoML/AutoMLExperiment/TunerFactory.cs index 0e89da262a..6acc5f37e6 100644 --- a/src/Microsoft.ML.AutoML/AutoMLExperiment/TunerFactory.cs +++ b/src/Microsoft.ML.AutoML/AutoMLExperiment/TunerFactory.cs @@ -7,7 +7,10 @@ namespace Microsoft.ML.AutoML { - internal interface ITunerFactory + /// + /// interface for all tuner factories. + /// + public interface ITunerFactory { ITuner CreateTuner(TrialSettings settings); } @@ -26,7 +29,7 @@ public ITuner CreateTuner(TrialSettings settings) var experimentSetting = _provider.GetService(); var searchSpace = settings.Pipeline.SearchSpace; var initParameter = settings.Pipeline.Parameter; - var isMaximize = experimentSetting.EvaluateMetric.IsMaximize; + var isMaximize = experimentSetting.IsMaximizeMetric; return new CostFrugalTuner(searchSpace, initParameter, !isMaximize); } diff --git a/src/Microsoft.ML.AutoML/SweepableEstimator/Estimator.cs b/src/Microsoft.ML.AutoML/SweepableEstimator/Estimator.cs index f7938bfe5b..bcb5da9b26 100644 --- a/src/Microsoft.ML.AutoML/SweepableEstimator/Estimator.cs +++ b/src/Microsoft.ML.AutoML/SweepableEstimator/Estimator.cs @@ -7,7 +7,7 @@ namespace Microsoft.ML.AutoML { - internal class Estimator + public class Estimator { protected Estimator() { diff --git a/src/Microsoft.ML.AutoML/SweepableEstimator/MultiModelPipeline.cs b/src/Microsoft.ML.AutoML/SweepableEstimator/MultiModelPipeline.cs index e11202ed98..eb1b58d1a1 100644 --- a/src/Microsoft.ML.AutoML/SweepableEstimator/MultiModelPipeline.cs +++ b/src/Microsoft.ML.AutoML/SweepableEstimator/MultiModelPipeline.cs @@ -10,7 +10,7 @@ namespace Microsoft.ML.AutoML { [JsonConverter(typeof(MultiModelPipelineConverter))] - internal class MultiModelPipeline + public class MultiModelPipeline { private static readonly StringEntity _nilStringEntity = new StringEntity("Nil"); private static readonly EstimatorEntity _nilSweepableEntity = new EstimatorEntity(null); diff --git a/src/Microsoft.ML.AutoML/SweepableEstimator/SweepableEstimator.cs b/src/Microsoft.ML.AutoML/SweepableEstimator/SweepableEstimator.cs index 26f3e8ef3b..35a6e0f3a8 100644 --- a/src/Microsoft.ML.AutoML/SweepableEstimator/SweepableEstimator.cs +++ b/src/Microsoft.ML.AutoML/SweepableEstimator/SweepableEstimator.cs @@ -14,7 +14,7 @@ namespace Microsoft.ML.AutoML /// Estimator with search space. /// [JsonConverter(typeof(SweepableEstimatorConverter))] - internal class SweepableEstimator : Estimator + public class SweepableEstimator : Estimator { private readonly Func> _factory; diff --git a/src/Microsoft.ML.AutoML/SweepableEstimator/SweepableEstimatorPipeline.cs b/src/Microsoft.ML.AutoML/SweepableEstimator/SweepableEstimatorPipeline.cs index 76e5682ba5..f04b7ae63b 100644 --- a/src/Microsoft.ML.AutoML/SweepableEstimator/SweepableEstimatorPipeline.cs +++ b/src/Microsoft.ML.AutoML/SweepableEstimator/SweepableEstimatorPipeline.cs @@ -11,7 +11,7 @@ namespace Microsoft.ML.AutoML { [JsonConverter(typeof(SweepableEstimatorPipelineConverter))] - internal class SweepableEstimatorPipeline + public class SweepableEstimatorPipeline { private readonly List _estimators; diff --git a/src/Microsoft.ML.AutoML/Tuner/ITuner.cs b/src/Microsoft.ML.AutoML/Tuner/ITuner.cs index 5844232c33..db522c192f 100644 --- a/src/Microsoft.ML.AutoML/Tuner/ITuner.cs +++ b/src/Microsoft.ML.AutoML/Tuner/ITuner.cs @@ -6,7 +6,7 @@ namespace Microsoft.ML.AutoML { - internal interface ITuner + public interface ITuner { Parameter Propose(TrialSettings settings);