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);