Fix the previous issues

This commit is contained in:
Wenbo Cao 2023-09-01 16:02:53 -05:00
parent 33803beb51
commit ea8ea43353
3 changed files with 35 additions and 18 deletions

View file

@ -12,19 +12,19 @@ using Microsoft.Extensions.DependencyInjection;
namespace BotSharp.Plugin.RoutingSpeeder.Controllers; namespace BotSharp.Plugin.RoutingSpeeder.Controllers;
[AllowAnonymous] [AllowAnonymous]
public class TrainIntentClassifierController : ControllerBase public class RoutingSpeederController : ControllerBase
{ {
private readonly IServiceProvider _service; private readonly IServiceProvider _service;
public TrainIntentClassifierController(IServiceProvider service) public RoutingSpeederController(IServiceProvider service)
{ {
_service = service; _service = service;
} }
[HttpPost("/intent/classifier/training")] [HttpPost("/routingspeeder/classifier/train")]
public IActionResult TrainIntentClassifier(TrainingParams trainingParams) public IActionResult TrainIntentClassifier(TrainingParams trainingParams)
{ {
var intentClassifier = _service.GetRequiredService<IntentClassifier>(); var intentClassifier = _service.GetRequiredService<IntentClassifier>();
intentClassifier.InitClassifer(trainingParams.Reference); intentClassifier.InitClassifer(trainingParams.Inference);
intentClassifier.Train(trainingParams); intentClassifier.Train(trainingParams);
return Ok(intentClassifier.Labels); return Ok(intentClassifier.Labels);
} }

View file

@ -11,35 +11,46 @@ using Tensorflow.Keras.Callbacks;
using System.Text.RegularExpressions; using System.Text.RegularExpressions;
using BotSharp.Plugin.RoutingSpeeder.Settings; using BotSharp.Plugin.RoutingSpeeder.Settings;
using BotSharp.Abstraction.MLTasks; using BotSharp.Abstraction.MLTasks;
using BotSharp.Abstraction.Knowledges.Settings;
using BotSharp.Plugin.RoutingSpeeder.Providers.Models; using BotSharp.Plugin.RoutingSpeeder.Providers.Models;
using Microsoft.Extensions.DependencyInjection; using Microsoft.Extensions.DependencyInjection;
using System.Linq; using System.Linq;
using Tensorflow.Keras; using Tensorflow.Keras;
using BotSharp.Abstraction.Knowledges.Settings;
using System.Numerics; using System.Numerics;
using Newtonsoft.Json; using Newtonsoft.Json;
using Tensorflow.Keras.Layers; using Tensorflow.Keras.Layers;
using BotSharp.Abstraction.Agents; using BotSharp.Abstraction.Agents;
using BotSharp.Abstraction.Knowledges;
namespace BotSharp.Plugin.RoutingSpeeder.Providers; namespace BotSharp.Plugin.RoutingSpeeder.Providers;
public class IntentClassifier public class IntentClassifier
{ {
private readonly IServiceProvider _services; private readonly IServiceProvider _services;
private KnowledgeBaseSettings _knowledgeBaseSettings;
Model _model; Model _model;
public Model model => _model; public Model model => _model;
private bool _isModelReady; private bool _isModelReady;
public bool isModelReady => _isModelReady; public bool isModelReady => _isModelReady;
private ClassifierSetting _settings; private ClassifierSetting _settings;
private string[] _labels => GetLabels(); private string[] _labels;
public string[] Labels => _labels; public string[] Labels => GetLabels();
public IntentClassifier(IServiceProvider services, ClassifierSetting settings) private int _numLabels
{
get
{
return Labels.Length;
}
}
public IntentClassifier(IServiceProvider services, ClassifierSetting settings, KnowledgeBaseSettings knowledgeBaseSettings)
{ {
_services = services; _services = services;
_settings = settings; _settings = settings;
_knowledgeBaseSettings = knowledgeBaseSettings;
} }
private void Reset() private void Reset()
@ -55,14 +66,15 @@ public class IntentClassifier
return; return;
} }
var vector = _services.GetRequiredService<ITextEmbedding>(); var vector = _services.GetServices<ITextEmbedding>()
.FirstOrDefault(x => x.GetType().FullName.EndsWith(_knowledgeBaseSettings.TextEmbedding));
var layers = new List<ILayer> var layers = new List<ILayer>
{ {
keras.layers.InputLayer((vector.Dimension), name: "Input"), keras.layers.InputLayer((vector.Dimension), name: "Input"),
keras.layers.Dense(256, activation:"relu"), keras.layers.Dense(256, activation:"relu"),
keras.layers.Dense(256, activation:"relu"), keras.layers.Dense(256, activation:"relu"),
keras.layers.Dense(_labels.Length, activation: keras.activations.Softmax) keras.layers.Dense(_numLabels, activation: keras.activations.Softmax)
}; };
_model = keras.Sequential(layers); _model = keras.Sequential(layers);
@ -92,7 +104,7 @@ public class IntentClassifier
var callbacks = new List<ICallback>() { earlyStop }; var callbacks = new List<ICallback>() { earlyStop };
var weights = LoadWeights(trainingParams.Reference); var weights = LoadWeights(trainingParams.Inference);
_model.fit(x, y, _model.fit(x, y,
batch_size: trainingParams.BatchSize, batch_size: trainingParams.BatchSize,
@ -159,12 +171,12 @@ public class IntentClassifier
{ {
var texts = File.ReadAllLines(filePath, Encoding.UTF8).Select(x => TextClean(x)).ToList(); var texts = File.ReadAllLines(filePath, Encoding.UTF8).Select(x => TextClean(x)).ToList();
vectorList.AddRange(vector.GetVectors(texts)); vectorList.AddRange(vector.GetVectors(texts));
string fileName = Path.GetFileNameWithoutExtension(filePath).Replace("intent.", ""); string fileName = Path.GetFileNameWithoutExtension(filePath);
labelList.AddRange(Enumerable.Repeat(fileName, texts.Count).ToList()); labelList.AddRange(Enumerable.Repeat(fileName, texts.Count).ToList());
} }
// Write label into local file // Write label into local file
var uniqueLabelList = labelList.Distinct().Select(x => x.Replace("intent.", "")).OrderBy(x => x).ToArray(); var uniqueLabelList = labelList.Distinct().OrderBy(x => x).ToArray();
File.WriteAllLines(saveLabelDirectory, uniqueLabelList); File.WriteAllLines(saveLabelDirectory, uniqueLabelList);
var x = np.zeros((vectorList.Count, vector.Dimension), dtype: np.float32); var x = np.zeros((vectorList.Count, vector.Dimension), dtype: np.float32);
@ -188,10 +200,15 @@ public class IntentClassifier
public string[] GetLabels() public string[] GetLabels()
{ {
var agentService = _services.CreateScope().ServiceProvider.GetRequiredService<IAgentService>(); if (_labels == null)
string rootDirectory = Path.Combine(agentService.GetDataDir(), _settings.MODEL_DIR, _settings.LABEL_FILE_NAME); {
var labelText = File.ReadAllLines(rootDirectory); var agentService = _services.CreateScope().ServiceProvider.GetRequiredService<IAgentService>();
return labelText.Select(x => x.Replace("intent.", "")).OrderBy(x => x).ToArray(); string rootDirectory = Path.Combine(agentService.GetDataDir(), _settings.MODEL_DIR, _settings.LABEL_FILE_NAME);
var labelText = File.ReadAllLines(rootDirectory);
_labels = labelText.OrderBy(x => x).ToArray();
}
return _labels;
} }
public string TextClean(string text) public string TextClean(string text)

View file

@ -10,5 +10,5 @@ public class TrainingParams
public int Epochs { get; set; } = 10; public int Epochs { get; set; } = 10;
public int BatchSize { get; set; } = 16; public int BatchSize { get; set; } = 16;
public float LearningRate { get; set; } = 1.0e-4f; public float LearningRate { get; set; } = 1.0e-4f;
public bool Reference { get; set; } = false; public bool Inference { get; set; } = false;
} }