diff --git a/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/KnowledgeBaseSettings.cs b/src/Infrastructure/BotSharp.Abstraction/Knowledges/Settings/KnowledgeBaseSettings.cs similarity index 81% rename from src/Infrastructure/BotSharp.Core/Plugins/Knowledges/KnowledgeBaseSettings.cs rename to src/Infrastructure/BotSharp.Abstraction/Knowledges/Settings/KnowledgeBaseSettings.cs index 0cb08a48..97f7f55c 100644 --- a/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/KnowledgeBaseSettings.cs +++ b/src/Infrastructure/BotSharp.Abstraction/Knowledges/Settings/KnowledgeBaseSettings.cs @@ -1,4 +1,4 @@ -namespace BotSharp.Core.Plugins.Knowledges; +namespace BotSharp.Abstraction.Knowledges.Settings; public class KnowledgeBaseSettings { diff --git a/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/KnowledgeBasePlugin.cs b/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/KnowledgeBasePlugin.cs index e0c74e83..8f862ca4 100644 --- a/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/KnowledgeBasePlugin.cs +++ b/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/KnowledgeBasePlugin.cs @@ -1,3 +1,4 @@ +using BotSharp.Abstraction.Knowledges.Settings; using BotSharp.Core.Plugins.Knowledges.Services; using Microsoft.Extensions.Configuration; diff --git a/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/Services/KnowledgeService.cs b/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/Services/KnowledgeService.cs index 9ed87e6a..3a627ec9 100644 --- a/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/Services/KnowledgeService.cs +++ b/src/Infrastructure/BotSharp.Core/Plugins/Knowledges/Services/KnowledgeService.cs @@ -1,7 +1,7 @@ using BotSharp.Abstraction.Knowledges.Models; +using BotSharp.Abstraction.Knowledges.Settings; using BotSharp.Abstraction.MLTasks; using BotSharp.Abstraction.VectorStorage; -using System.Text.Json; namespace BotSharp.Core.Plugins.Knowledges.Services; diff --git a/src/Infrastructure/BotSharp.Core/Templating/ResponseTemplateService.cs b/src/Infrastructure/BotSharp.Core/Templating/ResponseTemplateService.cs index b7be790a..ef2c23ba 100644 --- a/src/Infrastructure/BotSharp.Core/Templating/ResponseTemplateService.cs +++ b/src/Infrastructure/BotSharp.Core/Templating/ResponseTemplateService.cs @@ -46,6 +46,10 @@ public class ResponseTemplateService : IResponseTemplateService // Find response template var agentService = _services.GetRequiredService(); var dir = Path.Combine(agentService.GetAgentDataDir(agentId), "responses"); + if (!Directory.Exists(dir)) + { + return string.Empty; + } var responses = Directory.GetFiles(dir) .Where(f => f.Split(Path.DirectorySeparatorChar).Last().Split('.')[1] == message.IntentName) .ToList(); @@ -62,8 +66,15 @@ public class ResponseTemplateService : IResponseTemplateService // Convert args and execute data to dictionary var dict = new Dictionary(); - ExtractArgs(JsonSerializer.Deserialize(message.FunctionArgs), dict); - ExtractExecuteData(message.ExecutionData, dict); + if (!string.IsNullOrEmpty(message.FunctionArgs)) + { + ExtractArgs(JsonSerializer.Deserialize(message.FunctionArgs), dict); + } + + if (message.ExecutionData != null) + { + ExtractExecuteData(message.ExecutionData, dict); + } var text = render.Render(template, dict); diff --git a/src/Infrastructure/BotSharp.OpenAPI/Controllers/KnowledgeController.cs b/src/Infrastructure/BotSharp.OpenAPI/Controllers/KnowledgeController.cs index 3824f476..bea5597b 100644 --- a/src/Infrastructure/BotSharp.OpenAPI/Controllers/KnowledgeController.cs +++ b/src/Infrastructure/BotSharp.OpenAPI/Controllers/KnowledgeController.cs @@ -4,7 +4,7 @@ using Microsoft.AspNetCore.Http; using UglyToad.PdfPig.Content; using UglyToad.PdfPig; using BotSharp.Core.Plugins.Knowledges; - +using BotSharp.Abstraction.Knowledges.Settings; namespace BotSharp.OpenAPI.Controllers; diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/BotSharp.Plugin.RoutingSpeeder.csproj b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/BotSharp.Plugin.RoutingSpeeder.csproj index 79a54b27..fb145e8d 100644 --- a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/BotSharp.Plugin.RoutingSpeeder.csproj +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/BotSharp.Plugin.RoutingSpeeder.csproj @@ -7,6 +7,10 @@ 0.11.0 + + + + diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/IntentClassifier.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/IntentClassifier.cs new file mode 100644 index 00000000..382b3a98 --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/IntentClassifier.cs @@ -0,0 +1,245 @@ +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; +using BotSharp.Abstraction.Knowledges.Settings; +using System.Numerics; +using Newtonsoft.Json; +using Tensorflow.Keras.Layers; +using BotSharp.Abstraction.Agents; + +namespace BotSharp.Plugin.RoutingSpeeder.Providers; + +public class IntentClassifier +{ + private readonly IServiceProvider _services; + Model _model; + public Model model => _model; + private bool _isModelReady; + public bool isModelReady => _isModelReady; + private ClassifierSetting _settings; + + public IntentClassifier(IServiceProvider services, ClassifierSetting settings) + { + _services = services; + _settings = settings; + } + + private void Reset() + { + keras.backend.clear_session(); + _isModelReady = false; + } + + private void Build() + { + if (_isModelReady) + { + return; + } + + var vector = _services.GetRequiredService(); + + var layers = new List + { + keras.layers.InputLayer((vector.Dimension), name: "Input"), + keras.layers.Dense(256, activation:"relu"), + keras.layers.Dense(256, activation:"relu"), + keras.layers.Dense(GetLabels().Length, 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) + { + _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 agentService = _services.CreateScope().ServiceProvider.GetRequiredService(); + + var weightsFile = Path.Combine(agentService.GetDataDir(), _settings.MODEL_DIR, $"intent-classifier.h5"); + if (File.Exists(weightsFile)) + { + _model.load_weights(weightsFile); + _isModelReady = true; + Console.WriteLine($"Successfully load the weights!"); + } + else + { + Console.WriteLine("No available weights."); + } + return weightsFile; + } + + public (NDArray x, NDArray y) Vectorize(List items) + { + var vector = _services.GetRequiredService(); + + var x = np.zeros((items.Count, vector.Dimension), dtype: np.float32); + var y = np.zeros((items.Count, 1), dtype: np.float32); + + 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 NDArray GetTextEmbedding(string text) + { + var knowledgeSettings = _services.GetRequiredService(); + var embedding = _services.GetServices() + .FirstOrDefault(x => x.GetType().FullName.EndsWith(knowledgeSettings.TextEmbedding)); + + var x = np.zeros((1, embedding.Dimension), dtype: np.float32); + x[0] = embedding.GetVector(text); + return x; + } + + public (NDArray, NDArray) PrepareLoadData() + { + var agentService = _services.CreateScope().ServiceProvider.GetRequiredService(); + string rootDirectory = Path.Combine(agentService.GetDataDir(), _settings.RAW_DATA_DIR); + + + if (!Directory.Exists(rootDirectory)) + { + throw new Exception($"No training data found! Please put training data in this path: {rootDirectory}"); + } + + var vector = _services.GetRequiredService(); + + + var vectorList = new List(); + + var labelList = new List(); + foreach (var filePath in GetFiles()) + { + var texts = File.ReadAllLines(filePath, Encoding.UTF8).Select(x => TextClean(x)).ToList(); + vectorList.AddRange(vector.GetVectors(texts)); + string fileName = Path.GetFileNameWithoutExtension(filePath); + labelList.AddRange(Enumerable.Repeat(fileName, texts.Count).ToList()); + } + + var uniqueLabelList = labelList.Distinct().ToList(); + + var x = np.zeros((vectorList.Count, vector.Dimension), dtype: np.float32); + var y = np.zeros((vectorList.Count, 1), dtype: np.float32); + + for (int i = 0; i < vectorList.Count; i++) + { + x[i] = vectorList[i]; + y[i] = (float)uniqueLabelList.IndexOf(labelList[i]); + } + return (x, y); + } + + public string[] GetFiles() + { + var agentService = _services.CreateScope().ServiceProvider.GetRequiredService(); + string rootDirectory = Path.Combine(agentService.GetDataDir(), _settings.RAW_DATA_DIR); + return Directory.GetFiles(rootDirectory).OrderBy(x => x).ToArray(); + } + + public string[] GetLabels() + { + return GetFiles().Select(x => Path.GetFileNameWithoutExtension(x)).ToArray(); + } + + 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; + } + + public string Predict(NDArray vector) + { + if (!_isModelReady) + { + InitClassifer(); + } + + var prob = _model.predict(vector); + var probLabel = tf.arg_max(prob, -1).numpy(); + // var prediction = _settings.LabelMappingDict.First(x => x.Value == probLabel[0]).Key; + + var prediction = GetLabels()[probLabel[0]]; + // var prediction = GetLabels().Where((x, i) => i == probLabel[0]).First(); + + return prediction; + } + public void InitClassifer() + { + Reset(); + Build(); + LoadWeights(); + } + + public void Train() + { + var trainingParams = new TrainingParams(); + Reset(); + Build(); + (var x, var y) = PrepareLoadData(); + Fit(x, y, trainingParams); + + } +} 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/Models/TrainingParams.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/Models/TrainingParams.cs new file mode 100644 index 00000000..f3c822ac --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/Models/TrainingParams.cs @@ -0,0 +1,13 @@ +using System; +using System.Collections.Generic; +using System.Text; + +namespace BotSharp.Plugin.RoutingSpeeder.Providers.Models; + +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; +} diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingConversationHook.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingConversationHook.cs index 1a894c1c..fbb63acc 100644 --- a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingConversationHook.cs +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingConversationHook.cs @@ -2,25 +2,37 @@ using BotSharp.Abstraction.Agents.Enums; using BotSharp.Abstraction.Agents.Models; using BotSharp.Abstraction.Conversations; using BotSharp.Abstraction.Conversations.Models; -using BotSharp.Abstraction.Templating; +using BotSharp.Abstraction.MLTasks; using Microsoft.Extensions.DependencyInjection; using System; +using System.Linq; using System.Threading.Tasks; +using BotSharp.Plugin.RoutingSpeeder.Settings; +using BotSharp.Abstraction.Templating; +using BotSharp.Plugin.RoutingSpeeder.Providers; +using System.Runtime.InteropServices; namespace BotSharp.Plugin.RoutingSpeeder; public class RoutingConversationHook: ConversationHookBase { private readonly IServiceProvider _services; - public RoutingConversationHook(IServiceProvider services) + private RouterSpeederSettings _settings; + public RoutingConversationHook(IServiceProvider service, RouterSpeederSettings settings) { - _services = services; + _services = service; + _settings = settings; } - public override async Task BeforeCompletion(RoleDialogModel message) { + var intentClassifier = _services.GetRequiredService(); + var vector = intentClassifier.GetTextEmbedding(message.Content); + + // intentClassifier.Train(); // Utilize local discriminative model to predict intent - message.IntentName = "greeting"; + var predText = intentClassifier.Predict(vector); + + message.IntentName = predText; // Render by template var templateService = _services.GetRequiredService(); diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingSpeederPlugin.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/RoutingSpeederPlugin.cs index 28c856ba..c3dac3a2 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,13 @@ public class RoutingSpeederPlugin : IBotSharpPlugin { public void RegisterDI(IServiceCollection services, IConfiguration config) { + var settings = new RouterSpeederSettings(); + config.Bind("RouterSpeeder", settings); + services.AddSingleton(x => settings); + + services.AddSingleton(); + 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..d051f2e7 --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/classifierSetting.cs @@ -0,0 +1,18 @@ +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} + }; + + public string RAW_DATA_DIR { get; set; } = "raw_data"; + public string MODEL_DIR { get; set; } = "models"; +} 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..fbedf581 --- /dev/null +++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Settings/routerSpeedSettings.cs @@ -0,0 +1,9 @@ +using System; +using System.Collections.Generic; +using System.Text; + +namespace BotSharp.Plugin.RoutingSpeeder.Settings; + +public class RouterSpeederSettings +{ +} diff --git a/src/WebStarter/appsettings.json b/src/WebStarter/appsettings.json index 42ef6319..e409738f 100644 --- a/src/WebStarter/appsettings.json +++ b/src/WebStarter/appsettings.json @@ -47,11 +47,14 @@ } }, - "MetaAi": { - "fastText": { - "ModelPath": "crawl-300d-2M-subword.bin" - } - }, + "MetaAi": { + "fastText": { + "ModelPath": "crawl-300d-2M-subword.bin" + } + }, + + "RoutingSpeeder": { + }, "MetaMessenger": { "Endpoint": "https://graph.facebook.com",