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;
+}