diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/BotSharp.Plugin.RoutingSpeeder.csproj b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/BotSharp.Plugin.RoutingSpeeder.csproj
index fb145e8d..07118fd5 100644
--- a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/BotSharp.Plugin.RoutingSpeeder.csproj
+++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/BotSharp.Plugin.RoutingSpeeder.csproj
@@ -8,6 +8,7 @@
+
diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Controllers/TrainIntentClassifierController.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Controllers/TrainIntentClassifierController.cs
new file mode 100644
index 00000000..b4f5f18e
--- /dev/null
+++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Controllers/TrainIntentClassifierController.cs
@@ -0,0 +1,32 @@
+using System;
+using System.Collections.Generic;
+using System.Diagnostics;
+using System.Text;
+using System.Threading.Tasks;
+using BotSharp.Plugin.RoutingSpeeder.Providers;
+using BotSharp.Plugin.RoutingSpeeder.Providers.Models;
+using Microsoft.AspNetCore.Authorization;
+using Microsoft.AspNetCore.Mvc;
+using Microsoft.Extensions.DependencyInjection;
+
+namespace BotSharp.Plugin.RoutingSpeeder.Controllers;
+
+[AllowAnonymous]
+public class TrainIntentClassifierController : ControllerBase
+{
+ private readonly IServiceProvider _service;
+ public TrainIntentClassifierController(IServiceProvider service)
+ {
+ _service = service;
+ }
+
+ [HttpPost("/intent/classifier/training")]
+ public IActionResult TrainIntentClassifier(TrainingParams trainingParams)
+ {
+ var intentClassifier = _service.GetRequiredService();
+ intentClassifier.InitClassifer(trainingParams.Reference);
+ intentClassifier.Train(trainingParams);
+ return Ok(intentClassifier.Labels);
+ }
+
+}
diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/IntentClassifier.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/IntentClassifier.cs
index 3c3a33f2..5da02130 100644
--- a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/IntentClassifier.cs
+++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/IntentClassifier.cs
@@ -32,6 +32,10 @@ public class IntentClassifier
public bool isModelReady => _isModelReady;
private ClassifierSetting _settings;
+ private string[] _labels => GetLabels();
+
+ public string[] Labels => _labels;
+
public IntentClassifier(IServiceProvider services, ClassifierSetting settings)
{
_services = services;
@@ -50,17 +54,15 @@ public class IntentClassifier
{
return;
}
-
- var vector = _services.GetRequiredService();
- var labels = GetLabels();
+ 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(labels.Length, activation: keras.activations.Softmax)
+ keras.layers.Dense(_labels.Length, activation: keras.activations.Softmax)
};
_model = keras.Sequential(layers);
@@ -90,7 +92,7 @@ public class IntentClassifier
var callbacks = new List() { earlyStop };
- var weights = LoadWeights();
+ var weights = LoadWeights(trainingParams.Reference);
_model.fit(x, y,
batch_size: trainingParams.BatchSize,
@@ -104,42 +106,27 @@ public class IntentClassifier
_isModelReady = true;
}
- public string LoadWeights()
+ public string LoadWeights(bool inference = true)
{
var agentService = _services.CreateScope().ServiceProvider.GetRequiredService();
var weightsFile = Path.Combine(agentService.GetDataDir(), _settings.MODEL_DIR, $"intent-classifier.h5");
- if (File.Exists(weightsFile))
+
+ if (File.Exists(weightsFile) && inference)
{
_model.load_weights(weightsFile);
_isModelReady = true;
Console.WriteLine($"Successfully load the weights!");
+
}
else
{
- Console.WriteLine("No available weights.");
+ var logInfo = inference ? "No available weights." : "Will implement model training process and write trained weights into local";
+ Console.WriteLine(logInfo);
}
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();
@@ -164,20 +151,20 @@ public class IntentClassifier
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);
+ string fileName = Path.GetFileNameWithoutExtension(filePath).Replace("intent.", "");
labelList.AddRange(Enumerable.Repeat(fileName, texts.Count).ToList());
}
// Write label into local file
- var uniqueLabelList = labelList.Distinct().OrderBy(x => x).ToArray();
+ var uniqueLabelList = labelList.Distinct().Select(x => x.Replace("intent.", "")).OrderBy(x => x).ToArray();
File.WriteAllLines(saveLabelDirectory, uniqueLabelList);
var x = np.zeros((vectorList.Count, vector.Dimension), dtype: np.float32);
@@ -192,11 +179,11 @@ public class IntentClassifier
return (x, y);
}
- public string[] GetFiles()
+ public string[] GetFiles(string prefix = "intent")
{
var agentService = _services.CreateScope().ServiceProvider.GetRequiredService();
string rootDirectory = Path.Combine(agentService.GetDataDir(), _settings.RAW_DATA_DIR);
- return Directory.GetFiles(rootDirectory).OrderBy(x => x).ToArray();
+ return Directory.GetFiles(rootDirectory).Where(x => Path.GetFileNameWithoutExtension(x).StartsWith(prefix)).OrderBy(x => x).ToArray();
}
public string[] GetLabels()
@@ -204,7 +191,7 @@ public class IntentClassifier
var agentService = _services.CreateScope().ServiceProvider.GetRequiredService();
string rootDirectory = Path.Combine(agentService.GetDataDir(), _settings.MODEL_DIR, _settings.LABEL_FILE_NAME);
var labelText = File.ReadAllLines(rootDirectory);
- return labelText.OrderBy(x => x).ToArray();
+ return labelText.Select(x => x.Replace("intent.", "")).OrderBy(x => x).ToArray();
}
public string TextClean(string text)
@@ -235,24 +222,22 @@ public class IntentClassifier
return string.Empty;
}
- var prediction = GetLabels()[probLabel[0]];
+ var prediction = _labels[probLabel[0]];
return prediction;
}
- public void InitClassifer()
+ public void InitClassifer(bool inference = true)
{
Reset();
Build();
- LoadWeights();
+ LoadWeights(inference);
}
- public void Train()
+ public void Train(TrainingParams trainingParams)
{
- var trainingParams = new TrainingParams();
Reset();
(var x, var y) = PrepareLoadData();
Build();
Fit(x, y, trainingParams);
-
}
}
diff --git a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/Models/TrainingParams.cs b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/Models/TrainingParams.cs
index f3c822ac..7a4a8e67 100644
--- a/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/Models/TrainingParams.cs
+++ b/src/Plugins/BotSharp.Plugin.RoutingSpeeder/Providers/Models/TrainingParams.cs
@@ -10,4 +10,5 @@ public class TrainingParams
public int Epochs { get; set; } = 10;
public int BatchSize { get; set; } = 16;
public float LearningRate { get; set; } = 1.0e-4f;
+ public bool Reference { get; set; } = false;
}
diff --git a/src/WebStarter/data/models/intent-classifier.h5 b/src/WebStarter/data/models/intent-classifier.h5
index 13f2ebed..e4be7ee5 100644
Binary files a/src/WebStarter/data/models/intent-classifier.h5 and b/src/WebStarter/data/models/intent-classifier.h5 differ