Uh oh!
There was an error while loading. Please reload this page.
- Notifications
You must be signed in to change notification settings - Fork 1.9k
just a mere concept for seed, threads and logging#136
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Uh oh!
There was an error while loading. Please reload this page.
Changes from all commits
File filter
Filter by extension
Conversations
Uh oh!
There was an error while loading. Please reload this page.
Jump to
Uh oh!
There was an error while loading. Please reload this page.
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -11,6 +11,7 @@ | ||
| using System.Collections.Generic; | ||
| using System.Diagnostics; | ||
| using System.IO; | ||
| using static Microsoft.ML.Runtime.DefaultEnvironment; | ||
| namespace Microsoft.ML | ||
| { | ||
| @@ -48,15 +49,26 @@ public ScorerPipelineStep(Var<IDataView> data, Var<ITransformModel> model) | ||
| [DebuggerTypeProxy(typeof(LearningPipelineDebugProxy))] | ||
| public class LearningPipeline : ICollection<ILearningPipelineItem> | ||
| { | ||
| readonly internal IHostEnvironment Env; | ||
| private List<ILearningPipelineItem> Items { get; } = new List<ILearningPipelineItem>(); | ||
| /// <summary> | ||
| /// Construct an empty <see cref="LearningPipeline"/> object. | ||
| /// </summary> | ||
| public LearningPipeline() | ||
| public LearningPipeline(int? seed = null, int concurrency = 0) | ||
| ||
| { | ||
| var env = new DefaultEnvironment(seed: seed, conc: concurrency); | ||
| env.MessageRecieved += Env_MessageRecieved; | ||
| Env = env; | ||
| } | ||
| private void Env_MessageRecieved(object sender, ChannelMessageEventArgs e) | ||
| { | ||
| MessageOccured?.Invoke(this, e); | ||
| } | ||
| public event EventHandler<ChannelMessageEventArgs> MessageOccured; | ||
| ||
| /// <summary> | ||
| /// Get the count of ML components in the <see cref="LearningPipeline"/> object | ||
| /// </summary> | ||
| @@ -137,80 +149,76 @@ public PredictionModel<TInput, TOutput> Train<TInput, TOutput>() | ||
| where TInput : class | ||
| where TOutput : class, new() | ||
| { | ||
| Experiment experiment = Env.CreateExperiment(); | ||
| ILearningPipelineStep step = null; | ||
| List<ILearningPipelineLoader> loaders = new List<ILearningPipelineLoader>(); | ||
| List<Var<ITransformModel>> transformModels = new List<Var<ITransformModel>>(); | ||
| Var<ITransformModel> lastTransformModel = null; | ||
| using (var environment = new TlcEnvironment()) | ||
| foreach (ILearningPipelineItem currentItem in this) | ||
| { | ||
| Experiment experiment = environment.CreateExperiment(); | ||
| ILearningPipelineStep step = null; | ||
| List<ILearningPipelineLoader> loaders = new List<ILearningPipelineLoader>(); | ||
| List<Var<ITransformModel>> transformModels = new List<Var<ITransformModel>>(); | ||
| Var<ITransformModel> lastTransformModel = null; | ||
| if (currentItem is ILearningPipelineLoader loader) | ||
| loaders.Add(loader); | ||
| step = currentItem.ApplyStep(step, experiment); | ||
| if (step is ILearningPipelineDataStep dataStep && dataStep.Model != null) | ||
| transformModels.Add(dataStep.Model); | ||
| foreach (ILearningPipelineItem currentItem in this) | ||
| else if (step is ILearningPipelinePredictorStep predictorDataStep) | ||
| { | ||
| if (currentItem is ILearningPipelineLoader loader) | ||
| loaders.Add(loader); | ||
| step = currentItem.ApplyStep(step, experiment); | ||
| if (step is ILearningPipelineDataStep dataStep && dataStep.Model != null) | ||
| transformModels.Add(dataStep.Model); | ||
| else if (step is ILearningPipelinePredictorStep predictorDataStep) | ||
| if (lastTransformModel != null) | ||
| transformModels.Insert(0, lastTransformModel); | ||
| var localModelInput = new Transforms.ManyHeterogeneousModelCombiner | ||
| { | ||
| if (lastTransformModel != null) | ||
| transformModels.Insert(0, lastTransformModel); | ||
| var localModelInput = new Transforms.ManyHeterogeneousModelCombiner | ||
| { | ||
| PredictorModel = predictorDataStep.Model, | ||
| TransformModels = new ArrayVar<ITransformModel>(transformModels.ToArray()) | ||
| }; | ||
| var localModelOutput = experiment.Add(localModelInput); | ||
| var scorer = new Transforms.Scorer | ||
| { | ||
| PredictorModel = localModelOutput.PredictorModel | ||
| }; | ||
| var scorerOutput = experiment.Add(scorer); | ||
| lastTransformModel = scorerOutput.ScoringTransform; | ||
| step = new ScorerPipelineStep(scorerOutput.ScoredData, scorerOutput.ScoringTransform); | ||
| transformModels.Clear(); | ||
| } | ||
| } | ||
| PredictorModel = predictorDataStep.Model, | ||
| TransformModels = new ArrayVar<ITransformModel>(transformModels.ToArray()) | ||
| }; | ||
| if (transformModels.Count > 0) | ||
| { | ||
| transformModels.Insert(0,lastTransformModel); | ||
| var modelInput = new Transforms.ModelCombiner | ||
| var localModelOutput = experiment.Add(localModelInput); | ||
| var scorer = new Transforms.Scorer | ||
| { | ||
| Models = new ArrayVar<ITransformModel>(transformModels.ToArray()) | ||
| PredictorModel = localModelOutput.PredictorModel | ||
| }; | ||
| var modelOutput = experiment.Add(modelInput); | ||
| lastTransformModel = modelOutput.OutputModel; | ||
| var scorerOutput = experiment.Add(scorer); | ||
| lastTransformModel = scorerOutput.ScoringTransform; | ||
| step = new ScorerPipelineStep(scorerOutput.ScoredData, scorerOutput.ScoringTransform); | ||
| transformModels.Clear(); | ||
| } | ||
| } | ||
| experiment.Compile(); | ||
| foreach (ILearningPipelineLoader loader in loaders) | ||
| if (transformModels.Count > 0) | ||
| { | ||
| transformModels.Insert(0, lastTransformModel); | ||
| var modelInput = new Transforms.ModelCombiner | ||
| { | ||
| loader.SetInput(environment, experiment); | ||
| } | ||
| experiment.Run(); | ||
| Models = new ArrayVar<ITransformModel>(transformModels.ToArray()) | ||
| }; | ||
| ITransformModel model = experiment.GetOutput(lastTransformModel); | ||
| BatchPredictionEngine<TInput, TOutput> predictor; | ||
| using (var memoryStream = new MemoryStream()) | ||
| { | ||
| model.Save(environment, memoryStream); | ||
| var modelOutput = experiment.Add(modelInput); | ||
| lastTransformModel = modelOutput.OutputModel; | ||
| } | ||
| memoryStream.Position = 0; | ||
| experiment.Compile(); | ||
| foreach (ILearningPipelineLoader loader in loaders) | ||
| { | ||
| loader.SetInput(Env, experiment); | ||
| } | ||
| experiment.Run(); | ||
| predictor = environment.CreateBatchPredictionEngine<TInput, TOutput>(memoryStream); | ||
| ITransformModel model = experiment.GetOutput(lastTransformModel); | ||
| BatchPredictionEngine<TInput, TOutput> predictor; | ||
| using (var memoryStream = new MemoryStream()) | ||
| { | ||
| model.Save(Env, memoryStream); | ||
| return new PredictionModel<TInput, TOutput>(predictor, memoryStream); | ||
| } | ||
| memoryStream.Position = 0; | ||
| predictor = Env.CreateBatchPredictionEngine<TInput, TOutput>(memoryStream); | ||
| return new PredictionModel<TInput, TOutput>(predictor, memoryStream); | ||
| } | ||
| } | ||
| @@ -220,9 +228,9 @@ public PredictionModel<TInput, TOutput> Train<TInput, TOutput>() | ||
| /// <returns> | ||
| /// The IDataView that was returned by the pipeline. | ||
| /// </returns> | ||
| internal IDataView Execute(IHostEnvironment environment) | ||
| internal IDataView Execute() | ||
| { | ||
| Experiment experiment = environment.CreateExperiment(); | ||
| Experiment experiment = Env.CreateExperiment(); | ||
| ILearningPipelineStep step = null; | ||
| List<ILearningPipelineLoader> loaders = new List<ILearningPipelineLoader>(); | ||
| foreach (ILearningPipelineItem currentItem in this) | ||
| @@ -241,7 +249,7 @@ internal IDataView Execute(IHostEnvironment environment) | ||
| experiment.Compile(); | ||
| foreach (ILearningPipelineLoader loader in loaders) | ||
| { | ||
| loader.SetInput(environment, experiment); | ||
| loader.SetInput(Env, experiment); | ||
| } | ||
| experiment.Run(); | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,96 @@ | ||
| using Microsoft.ML.Runtime.Data; | ||
| using System; | ||
| namespace Microsoft.ML.Runtime | ||
| { | ||
| public sealed class DefaultEnvironment : HostEnvironmentBase<DefaultEnvironment> | ||
Member There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Is there a reason why this type is public? | ||
| { | ||
| public DefaultEnvironment(int? seed = null, int conc = 0) | ||
| : this(RandomUtils.Create(seed), true, conc) | ||
| { | ||
| } | ||
| public DefaultEnvironment(IRandom rand, bool verbose, int conc, string shortName = null, string parentFullName = null) : base(rand, verbose, conc, shortName, parentFullName) | ||
| { | ||
| EnsureDispatcher<ChannelMessage>(); | ||
| AddListener<ChannelMessage>(OnMessageRecieved); | ||
| } | ||
| void OnMessageRecieved(IMessageSource sender, ChannelMessage msg) | ||
| { | ||
| ChannelMessageEventArgs eventArgs = new ChannelMessageEventArgs() { Message = msg }; | ||
| MessageRecieved?.Invoke(this, eventArgs); | ||
| } | ||
| public event EventHandler<ChannelMessageEventArgs> MessageRecieved; | ||
| public class ChannelMessageEventArgs : EventArgs | ||
| { | ||
| public ChannelMessage Message { get; set; } | ||
| } | ||
| private sealed class Channel : ChannelBase | ||
| { | ||
| public Channel(DefaultEnvironment master, ChannelProviderBase parent, string shortName, Action<IMessageSource, ChannelMessage> dispatch) | ||
| : base(master, parent, shortName, dispatch) | ||
| { | ||
| } | ||
| } | ||
| private sealed class Host : HostBase | ||
| { | ||
| public new bool IsCancelled => Root.IsCancelled; | ||
| public Host(HostEnvironmentBase<DefaultEnvironment> source, string shortName, string parentFullName, IRandom rand, bool verbose, int? conc) | ||
| : base(source, shortName, parentFullName, rand, verbose, conc) | ||
| { | ||
| } | ||
| protected override IChannel CreateCommChannel(ChannelProviderBase parent, string name) | ||
| { | ||
| Contracts.AssertValue(parent); | ||
| Contracts.Assert(parent is Host); | ||
| Contracts.AssertNonEmpty(name); | ||
| return new Channel(Root, parent, name, GetDispatchDelegate<ChannelMessage>()); | ||
| } | ||
| protected override IPipe<TMessage> CreatePipe<TMessage>(ChannelProviderBase parent, string name) | ||
| { | ||
| Contracts.AssertValue(parent); | ||
| Contracts.Assert(parent is Host); | ||
| Contracts.AssertNonEmpty(name); | ||
| return new Pipe<TMessage>(parent, name, GetDispatchDelegate<TMessage>()); | ||
| } | ||
| protected override IHost RegisterCore(HostEnvironmentBase<DefaultEnvironment> source, string shortName, string parentFullName, IRandom rand, bool verbose, int? conc) | ||
| { | ||
| return new Host(source, shortName, parentFullName, rand, verbose, conc); | ||
| } | ||
| } | ||
| protected override IHost RegisterCore(HostEnvironmentBase<DefaultEnvironment> source, string shortName, string parentFullName, IRandom rand, bool verbose, int? conc) | ||
| { | ||
| Contracts.AssertValue(rand); | ||
| Contracts.AssertValueOrNull(parentFullName); | ||
| Contracts.AssertNonEmpty(shortName); | ||
| Contracts.Assert(source == this || source is Host); | ||
| return new Host(source, shortName, parentFullName, rand, verbose, conc); | ||
| } | ||
| protected override IChannel CreateCommChannel(ChannelProviderBase parent, string name) | ||
| { | ||
| Contracts.AssertValue(parent); | ||
| Contracts.Assert(parent is DefaultEnvironment); | ||
| Contracts.AssertNonEmpty(name); | ||
| return new Channel(this, parent, name, GetDispatchDelegate<ChannelMessage>()); | ||
| } | ||
| protected override IPipe<TMessage> CreatePipe<TMessage>(ChannelProviderBase parent, string name) | ||
| { | ||
| Contracts.AssertValue(parent); | ||
| Contracts.Assert(parent is DefaultEnvironment); | ||
| Contracts.AssertNonEmpty(name); | ||
| return new Pipe<TMessage>(parent, name, GetDispatchDelegate<TMessage>()); | ||
| } | ||
| } | ||
| } | ||
Uh oh!
There was an error while loading. Please reload this page.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
HostEnvironmentBase just happens to be IDisposable.
Since this design makes the pipeline own an instance of Env, your have to make the class disposable,
Alternatively, you can take an optional environment and then the problem of owning env becomse the callers problem.
If a user has to pass an environment to this class there are several potental benefits: