From 0d017f1e863c7809190b9ea7fe1ccabefee6f557 Mon Sep 17 00:00:00 2001 From: Wenbo Cao <104199@smsassist.com> Date: Thu, 31 Aug 2023 09:52:09 -0500 Subject: [PATCH] Add DialoguePrediction --- .../BotSharp.Plugin.RoutingSpeeder.csproj | 6 + .../Providers/DialogueClassifier.cs | 144 ++++++++++++++++++ .../Models/DialoguePredictionModel.cs | 13 ++ .../Providers/fastTextEmbeddingProvider.cs | 70 +++++++++ .../RoutingConversationHook.cs | 16 ++ .../RoutingSpeederPlugin.cs | 8 + .../Settings/classifierSetting.cs | 21 +++ .../Settings/fastTextSetting.cs | 7 + .../Settings/routerSpeedSettings.cs | 11 ++ .../Settings/trainingParams.cs | 13 ++ 10 files changed, 309 insertions(+) create mode 100644 src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/DialogueClassifier.cs create mode 100644 src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/Models/DialoguePredictionModel.cs create mode 100644 src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/fastTextEmbeddingProvider.cs create mode 100644 src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/classifierSetting.cs create mode 100644 src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/fastTextSetting.cs create mode 100644 src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/routerSpeedSettings.cs create mode 100644 src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/trainingParams.cs diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/BotSharp.Plugin.RoutingSpeeder.csproj b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/BotSharp.Plugin.RoutingSpeeder.csproj index 79a54b27..e0a8e7af 100644 --- a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/BotSharp.Plugin.RoutingSpeeder.csproj +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/BotSharp.Plugin.RoutingSpeeder.csproj @@ -11,4 +11,10 @@ + + + + + + diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/DialogueClassifier.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/DialogueClassifier.cs new file mode 100644 index 00000000..788eb641 --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/DialogueClassifier.cs @@ -0,0 +1,144 @@ +using System; +using System.IO; +using System.Text; +using System.Collections.Generic; +using Tensorflow; +using static Tensorflow.KerasApi; +using Tensorflow.Keras.Engine; +using Tensorflow.NumPy; +using static Tensorflow.Binding; +using Tensorflow.Keras.Callbacks; +using System.Text.RegularExpressions; +using BotSharp.Plugin.RoutingSpeeder.Settings; +using BotSharp.Abstraction.MLTasks; +using BotSharp.Plugin.RoutingSpeeder.Providers.Models; +using Microsoft.Extensions.DependencyInjection; +using System.Linq; +using Tensorflow.Keras; + +namespace BotSharp.Plugin.RoutingSpeeder.Providers; + +public class DialogueClassifier +{ + private readonly IServiceProvider _services; + Model _model; + public Model model => _model; + private bool _isModelReady; + public bool isModelReady => _isModelReady; + private classifierSetting _settings; + + public DialogueClassifier(IServiceProvider services, classifierSetting settings) + { + _services = services; + _settings = settings; + } + + private void Reset() + { + keras.backend.clear_session(); + _isModelReady = false; + } + + private void Build() + { + if (_isModelReady) + { + return; + } + + var layers = new List + { + keras.layers.InputLayer((300), name: "Input"), + keras.layers.Dense(256, activation:"relu"), + keras.layers.Dense(256, activation:"relu"), + keras.layers.Dense(_settings.labelMappingDict.Count, activation: keras.activations.Softmax) + }; + _model = keras.Sequential(layers); + +#if DEBUG + Console.WriteLine(); + _model.summary(); +#endif + _isModelReady = true; + } + + private void Fit(NDArray x, NDArray y, TrainingParams trainingParams) + { + // release more memory + var vector = _services.GetRequiredService(); + // vector.UnloadModel(); + + _model.compile(optimizer: keras.optimizers.Adam(trainingParams.LearningRate), + loss: keras.losses.SparseCategoricalCrossentropy(), + metrics: new[] { "accuracy" } + ); + + CallbackParams callback_parameters = new CallbackParams + { + Model = _model, + Epochs = trainingParams.Epochs, + Verbose = 1, + Steps = 10 + }; + + ICallback earlyStop = new EarlyStopping(callback_parameters, "accuracy"); + + var callbacks = new List() { earlyStop }; + + var weights = LoadWeights(); + + _model.fit(x, y, + batch_size: trainingParams.BatchSize, + epochs: trainingParams.Epochs, + callbacks: callbacks, + // validation_split: 0.1f, + shuffle: true); + + _model.save_weights(weights); + + _isModelReady = true; + } + + public string LoadWeights() + { + var weightsFile = Path.Combine(_settings.MODEL_DIR, $"wo-dialogue-classifier.h5"); + if (File.Exists(weightsFile)) + { + _model.load_weights(weightsFile); + Console.WriteLine($"Successfully load the weights!"); + } + else + { + Console.WriteLine("No available weights."); + } + return weightsFile; + } + + public (NDArray x, NDArray y) Vectorize(List items) + { + var x = np.zeros((items.Count, 300), dtype: np.float32); + var y = np.zeros((items.Count, 1), dtype: np.float32); + + var vector = _services.GetRequiredService(); + + for (int i = 0; i < items.Count; i++) + { + x[i] = vector.GetVector(TextClean(items[i].text)); + if (_settings.labelMappingDict.ContainsKey(items[i].label)) + { + y[i] = _settings.labelMappingDict[items[i].label]; + } + } + return (x, y); + } + public string TextClean(string text) + { + // Remove punctuation + // Remove digits + // To lowercase + var processedText = Regex.Replace(text, "[AB0-9]", " "); + processedText = string.Join("", processedText.Select(c => char.IsPunctuation(c) ? ' ' : c).ToList()); + processedText = processedText.Replace(" ", " ").ToLower(); + return processedText; + } +} diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/Models/DialoguePredictionModel.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/Models/DialoguePredictionModel.cs new file mode 100644 index 00000000..4641b9cd --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/Models/DialoguePredictionModel.cs @@ -0,0 +1,13 @@ +using System; +using System.Collections.Generic; +using System.Text; + +namespace BotSharp.Plugin.RoutingSpeeder.Providers.Models; + +public class DialoguePredictionModel +{ + public int Id { get; set; } + public string text { get; set; } + public string? label { get; set; } + public string? prediction { get; set; } +} diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/fastTextEmbeddingProvider.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/fastTextEmbeddingProvider.cs new file mode 100644 index 00000000..79925418 --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/fastTextEmbeddingProvider.cs @@ -0,0 +1,70 @@ +using System; +using System.Collections.Generic; +using System.IO; +using System.Runtime; +using System.Text; +using System.Text.RegularExpressions; +using BotSharp.Abstraction.MLTasks; +using BotSharp.Plugin.RoutingSpeeder.Settings; +using FastText.NetWrapper; + +namespace BotSharp.Plugin.RoutingSpeeder.Providers; + +public class fastTextEmbeddingProvider : ITextEmbedding +{ + private FastTextWrapper _fastText; + private readonly fastTextSetting _settings; + + public int Dimension + { + get + { + if (!_fastText.IsModelReady()) + { + _fastText.LoadModel(_settings.ModelPath); + } + return _fastText.GetModelDimension(); + } + } + + public fastTextEmbeddingProvider(fastTextSetting settings) + { + _settings = settings; + + } + + public float[] GetVector(string text) + { + LoadModel(); + return _fastText.GetSentenceVector(text); + } + + public List GetVectors(List texts) + { + LoadModel(); + var vectors = new List(); + for (int i = 0; i < texts.Count; i++) + { + vectors.Add(GetVector(texts[i])); + } + return vectors; + } + + private void LoadModel() + { + if (_fastText == null) + { + if (!File.Exists(_settings.ModelPath)) + { + throw new FileNotFoundException($"Can't load pre-trained word vectors from {_settings.ModelPath}.\n Try to download from https://fasttext.cc/docs/en/english-vectors.html."); + } + + _fastText = new FastTextWrapper(); + + if (!_fastText.IsModelReady()) + { + _fastText.LoadModel(_settings.ModelPath); + } + } + } +} diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingConversationHook.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingConversationHook.cs index 50d74fe8..bc37e117 100644 --- a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingConversationHook.cs +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingConversationHook.cs @@ -1,13 +1,29 @@ using BotSharp.Abstraction.Conversations; using BotSharp.Abstraction.Conversations.Models; +using BotSharp.Abstraction.MLTasks; +using Microsoft.Extensions.DependencyInjection; +using System; +using System.Linq; using System.Threading.Tasks; +using FastText.NetWrapper; +using BotSharp.Plugin.RoutingSpeeder.Settings; namespace BotSharp.Plugin.RoutingSpeeder; public class RoutingConversationHook: ConversationHookBase { + private readonly IServiceProvider _services; + private routerSpeedSettings _settings; + public RoutingConversationHook(IServiceProvider service, routerSpeedSettings settings) + { + _services = service; + _settings = settings; + } public override async Task BeforeCompletion(RoleDialogModel message) { + var embedding = _services.GetServices() + .FirstOrDefault(x => x.GetType().FullName.EndsWith(_settings.TextEmbedding)); + // Utilize local discriminative model to predict intent message.Content = "response content"; message.StopCompletion = true; diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingSpeederPlugin.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingSpeederPlugin.cs index 28c856ba..bcb91335 100644 --- a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingSpeederPlugin.cs +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingSpeederPlugin.cs @@ -1,5 +1,8 @@ using BotSharp.Abstraction.Conversations; +using BotSharp.Abstraction.MLTasks; using BotSharp.Abstraction.Plugins; +using BotSharp.Plugin.RoutingSpeeder.Settings; +using BotSharp.Plugin.RoutingSpeeder.Providers; using Microsoft.Extensions.Configuration; using Microsoft.Extensions.DependencyInjection; @@ -9,6 +12,11 @@ public class RoutingSpeederPlugin : IBotSharpPlugin { public void RegisterDI(IServiceCollection services, IConfiguration config) { + var settings = new routerSpeedSettings(); + config.Bind("routerSpeed", settings); + services.AddSingleton(x => settings); + services.AddSingleton(x => settings.fastText); services.AddScoped(); + services.AddSingleton(); } } diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/classifierSetting.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/classifierSetting.cs new file mode 100644 index 00000000..2fe13237 --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/classifierSetting.cs @@ -0,0 +1,21 @@ +using System; +using System.Collections.Generic; +using System.Text; + +namespace BotSharp.Plugin.RoutingSpeeder.Settings; + +public class classifierSetting +{ + + public Dictionary labelMappingDict { get; set; } = new Dictionary() + { + {"goodbye", 0f}, + {"greeting", 1f}, + {"other", 2f}, + {"wo-followup", 3f}, + {"wo-identifer", 4f}, + {"wo-scheduler", 5} + }; + public string RAW_DATA_DIR { get; set; } = "C:\\new_wenbocao\\one_brain\\WebStarter\\data\\raw_data"; + public string MODEL_DIR { get; set; } = "C:\\new_wenbocao\\one_brain\\WebStarter\\data\\models"; +} diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/fastTextSetting.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/fastTextSetting.cs new file mode 100644 index 00000000..e5554402 --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/fastTextSetting.cs @@ -0,0 +1,7 @@ + +namespace BotSharp.Plugin.RoutingSpeeder.Settings; + +public class fastTextSetting +{ + public string ModelPath { get; set; } +} diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/routerSpeedSettings.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/routerSpeedSettings.cs new file mode 100644 index 00000000..300ec9e0 --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/routerSpeedSettings.cs @@ -0,0 +1,11 @@ +using System; +using System.Collections.Generic; +using System.Text; + +namespace BotSharp.Plugin.RoutingSpeeder.Settings; + +public class routerSpeedSettings +{ + public fastTextSetting fastText { get; set; } + public string TextEmbedding { get; set; } +} diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/trainingParams.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/trainingParams.cs new file mode 100644 index 00000000..fcf15ce3 --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/trainingParams.cs @@ -0,0 +1,13 @@ +using System; +using System.Collections.Generic; +using System.Text; + +namespace BotSharp.Plugin.RoutingSpeeder.Settings; + +public class TrainingParams +{ + public int ClientId { get; set; } + public int Epochs { get; set; } = 10; + public int BatchSize { get; set; } = 16; + public float LearningRate { get; set; } = 1.0e-4f; +}