From ad309dc68d2eb543ce5052493e5984e45c8348fc Mon Sep 17 00:00:00 2001 From: "haiping008@gmail.com" Date: Mon, 18 Jun 2018 17:15:32 -0500 Subject: [PATCH] train model by input contexts. --- BotSharp.Core/Agents/AgentDriver.cs | 4 + BotSharp.Core/BotSharp.Core.csproj | 4 + .../Dialogflow/AgentImporterInDialogflow.cs | 6 + BotSharp.Core/Engines/Rasa/RasaAi.cs | 114 ++++++++---- .../Engines/Rasa/RasaIntentExpression.cs | 6 +- .../Engines/Rasa/RasaTraningEntity.cs | 1 + BotSharp.Core/Engines/RequestExtension.cs | 163 +++--------------- BotSharp.Core/Intents/ContextModelMapping.cs | 25 --- BotSharp.Core/Intents/Intent.cs | 4 +- BotSharp.UnitTest/AgentTest.cs | 4 +- BotSharp.UnitTest/BotSharp.UnitTest.csproj | 3 - BotSharp.UnitTest/IntentTest.cs | 23 +++ BotSharp.UnitTest/Settings/settings.bot.json | 2 +- 13 files changed, 147 insertions(+), 212 deletions(-) delete mode 100644 BotSharp.Core/Intents/ContextModelMapping.cs diff --git a/BotSharp.Core/Agents/AgentDriver.cs b/BotSharp.Core/Agents/AgentDriver.cs index 5e18394c..c100b2f8 100644 --- a/BotSharp.Core/Agents/AgentDriver.cs +++ b/BotSharp.Core/Agents/AgentDriver.cs @@ -93,6 +93,7 @@ namespace BotSharp.Core.Agents { Intent = intent.Name, Text = String.Join("", exp.Data.OrderBy(x => x.UpdatedTime).Select(x => x.Text)), + ContextHash = intent.ContextHash }; // convert entity format @@ -137,6 +138,9 @@ namespace BotSharp.Core.Agents }); }); + // remove Default Fallback Intent + trainingData.UserSays = trainingData.UserSays.Where(x => x.Intent != "Default Fallback Intent").ToList(); + return trainingData; } } diff --git a/BotSharp.Core/BotSharp.Core.csproj b/BotSharp.Core/BotSharp.Core.csproj index 9e372958..0e69260c 100644 --- a/BotSharp.Core/BotSharp.Core.csproj +++ b/BotSharp.Core/BotSharp.Core.csproj @@ -23,6 +23,10 @@ https://github.com/Oceania2018/BotSharp + + TRACE;DEBUG;MODEL_PER_CONTEXTS + + diff --git a/BotSharp.Core/Engines/Dialogflow/AgentImporterInDialogflow.cs b/BotSharp.Core/Engines/Dialogflow/AgentImporterInDialogflow.cs index 63e8f997..dc05f813 100644 --- a/BotSharp.Core/Engines/Dialogflow/AgentImporterInDialogflow.cs +++ b/BotSharp.Core/Engines/Dialogflow/AgentImporterInDialogflow.cs @@ -108,6 +108,12 @@ namespace BotSharp.Core.Engines { string expressionJson = File.ReadAllText($"{expressionFileName}"); intent.UserSays = JsonConvert.DeserializeObject>(expressionJson); + + // remove @sys.ignore + intent.UserSays.ForEach(say => + { + say.Data.Where(x => x.Meta == "@sys.ignore").ToList().ForEach(x => x.Meta = null); + }); } } diff --git a/BotSharp.Core/Engines/Rasa/RasaAi.cs b/BotSharp.Core/Engines/Rasa/RasaAi.cs index 80ffec3e..9856e5e6 100644 --- a/BotSharp.Core/Engines/Rasa/RasaAi.cs +++ b/BotSharp.Core/Engines/Rasa/RasaAi.cs @@ -2,6 +2,7 @@ using BotSharp.Core.Entities; using BotSharp.Core.Intents; using BotSharp.Core.Models; +using DotNetToolkit; using EntityFrameworkCore.BootKit; using Microsoft.EntityFrameworkCore; using Microsoft.Extensions.Configuration; @@ -49,9 +50,6 @@ namespace BotSharp.Core.Engines var corpus = agent.GrabCorpus(dc); - // remove Default Fallback Intent - corpus.UserSays = corpus.UserSays.Where(x => x.Intent != "Default Fallback Intent").ToList(); - string json = JsonConvert.SerializeObject(new { rasa_nlu_data = corpus }, new JsonSerializerSettings { @@ -59,14 +57,10 @@ namespace BotSharp.Core.Engines NullValueHandling = NullValueHandling.Ignore }); -#if RASA_NLU_0_11 - rest.AddParameter("application/json", json, ParameterType.RequestBody); -#else - string trainingConfig = agent.Language == "zh" ? "config_jieba_mitie_sklearn.yml" : "config_mitie_sklearn.yml"; + string trainingConfig = agent.Language == "zh" ? "config_jieba_mitie_sklearn.yml" : "config_spacy.yml"; string body = File.ReadAllText($"{Database.ContentRootPath}{Path.DirectorySeparatorChar}Settings{Path.DirectorySeparatorChar}{trainingConfig}"); body = $"{body}\r\ndata: {json}"; rest.AddParameter("application/x-yml", body, ParameterType.RequestBody); -#endif var response = client.Execute(rest); @@ -88,47 +82,93 @@ namespace BotSharp.Core.Engines } } - public string TrainWithContexts() + public void TrainWithContexts() { var corpus = agent.GrabCorpus(dc); - string json = JsonConvert.SerializeObject(new { rasa_nlu_data = corpus }, - new JsonSerializerSettings + var client = new RestClient($"{Database.Configuration.GetSection("Rasa:Nlu").Value}"); + + var contextHashs = corpus.UserSays + .Select(x => x.ContextHash) + .Distinct() + .ToList(); + + contextHashs.ForEach(ctx => + { + var data = new RasaTrainingData { - ContractResolver = new CamelCasePropertyNamesContractResolver(), - NullValueHandling = NullValueHandling.Ignore + Entities = corpus.Entities, + UserSays = corpus.UserSays.Where(x => x.ContextHash == ctx).ToList() + }; + + // meet minimal requirement + // at least 2 different classes + int count = data.UserSays + .Select(x => x.Intent) + .Distinct().Count(); + + if (count < 2) + { + data.UserSays.Add(new RasaIntentExpression + { + Intent = "Intent2", + Text = Guid.NewGuid().ToString("N") + }); + + data.UserSays.Add(new RasaIntentExpression + { + Intent = "Intent2", + Text = Guid.NewGuid().ToString("N") + }); + } + + // at least 2 corpus per intent + data.UserSays.Select(x => x.Intent) + .Distinct() + .ToList() + .ForEach(intent => + { + if(data.UserSays.Count(x => x.Intent == intent) < 2) + { + data.UserSays.Add(new RasaIntentExpression + { + Intent = intent, + Text = Guid.NewGuid().ToString("N") + }); + } }); - var client = new RestClient($"{Database.Configuration.GetSection("Rasa:Host").Value}"); - var rest = new RestRequest("train", Method.POST); - rest.AddQueryParameter("project", agent.Id); - rest.AddParameter("application/json", json, ParameterType.RequestBody); + string json = JsonConvert.SerializeObject(new { rasa_nlu_data = data }, + new JsonSerializerSettings + { + ContractResolver = new CamelCasePropertyNamesContractResolver(), + NullValueHandling = NullValueHandling.Ignore + }); - var response = client.Execute(rest); + var rest = new RestRequest("train", Method.POST); + rest.AddQueryParameter("project", agent.Id); + rest.AddQueryParameter("model", ctx); + string trainingConfig = agent.Language == "zh" ? "config_jieba_mitie_sklearn.yml" : "config_spacy.yml"; + string body = File.ReadAllText($"{Database.ContentRootPath}{Path.DirectorySeparatorChar}Settings{Path.DirectorySeparatorChar}{trainingConfig}"); + body = $"{body}\r\ndata: {json}"; + rest.AddParameter("application/x-yml", body, ParameterType.RequestBody); - if (response.IsSuccessful) - { - var result = JObject.Parse(response.Content); + var response = client.Execute(rest); - string modelName = result["info"].Value().Split(": ")[1]; - - dc.Table().Add(new ContextModelMapping + if (response.IsSuccessful) { - AgentId = agent.Id, - ModelName = modelName, - //ContextId = contextId - }); + var result = JObject.Parse(response.Content); - return modelName; - } - else - { - var result = JObject.Parse(response.Content); + string modelName = result["info"].Value().Split(": ")[1]; + } + else + { + var result = JObject.Parse(response.Content); + Console.WriteLine(result["error"]); + result["error"].Log(); + } + }); - Console.WriteLine(result["error"]); - - return String.Empty; - } } } } diff --git a/BotSharp.Core/Engines/Rasa/RasaIntentExpression.cs b/BotSharp.Core/Engines/Rasa/RasaIntentExpression.cs index 6ab8d0a3..b3dfcd4c 100644 --- a/BotSharp.Core/Engines/Rasa/RasaIntentExpression.cs +++ b/BotSharp.Core/Engines/Rasa/RasaIntentExpression.cs @@ -1,4 +1,5 @@ -using System; +using Newtonsoft.Json; +using System; using System.Collections.Generic; using System.Text; @@ -12,7 +13,10 @@ namespace BotSharp.Core.Models public String Text { get; set; } public String Intent { get; set; } + + [JsonIgnore] public String ContextHash { get; set; } + public List Entities { get; set; } } } diff --git a/BotSharp.Core/Engines/Rasa/RasaTraningEntity.cs b/BotSharp.Core/Engines/Rasa/RasaTraningEntity.cs index 81fcf50a..67a38a28 100644 --- a/BotSharp.Core/Engines/Rasa/RasaTraningEntity.cs +++ b/BotSharp.Core/Engines/Rasa/RasaTraningEntity.cs @@ -7,6 +7,7 @@ namespace BotSharp.Core.Adapters.Rasa { public class RasaTraningEntity { + [JsonIgnore] public String EntityType { get; set; } [JsonProperty("value")] diff --git a/BotSharp.Core/Engines/RequestExtension.cs b/BotSharp.Core/Engines/RequestExtension.cs index 4266c5ae..e7bfcc68 100644 --- a/BotSharp.Core/Engines/RequestExtension.cs +++ b/BotSharp.Core/Engines/RequestExtension.cs @@ -30,7 +30,12 @@ namespace BotSharp.Core.Engines AIResponse aiResponse = new AIResponse(); Database dc = rasa.dc; +#if MODEL_PER_CONTEXTS + string model = GetModelPerContexts(rasa, request); + var result = CallRasa(rasa.agent.Id, request.Query.First(), model); +#else var result = CallRasa(rasa.agent.Id, request.Query.First(), rasa.agent.Id); +#endif RasaResponse response = result.Data; aiResponse.Id = Guid.NewGuid().ToString(); aiResponse.Lang = rasa.agent.Language; @@ -114,7 +119,7 @@ namespace BotSharp.Core.Engines }; } - response.IntentRanking = response.IntentRanking.Where(x => x.Confidence > decimal.Parse("0.1")).ToList(); + response.IntentRanking = response.IntentRanking.Where(x => x.Confidence > decimal.Parse("0.3")).ToList(); response.IntentRanking = response.IntentRanking.Where(x => intents.Select(i => i.Name).Contains(x.Name)).ToList(); // add Default Fallback Intent @@ -165,10 +170,14 @@ namespace BotSharp.Core.Engines // convert to Standard entity value if (!String.IsNullOrEmpty(p.Value) && !p.DataType.StartsWith("@sys.")) { - p.Value = agent.Entities.FirstOrDefault(x => x.Name == p.Name).Entries.FirstOrDefault((entry) => { - return entry.Value.ToLower() == p.Value.ToLower() || - entry.Synonyms.Select(synonym => synonym.Synonym.ToLower()).Contains(p.Value.ToLower()); - })?.Value; + p.Value = agent.Entities + .FirstOrDefault(x => x.Name == p.DataType.Substring(1)) + .Entries + .FirstOrDefault((entry) => + { + return entry.Value.ToLower() == p.Value.ToLower() || + entry.Synonyms.Select(synonym => synonym.Synonym.ToLower()).Contains(p.Value.ToLower()); + })?.Value; } // fixed entity per request @@ -202,11 +211,14 @@ namespace BotSharp.Core.Engines } else { - msg.Speech = msg.Speech.StartsWith("[") ? + if (msg.Speech != "[]") + { + msg.Speech = msg.Speech.StartsWith("[") ? ArrayHelper.GetRandom(msg.Speech.Substring(2, msg.Speech.Length - 4).Split("\",\"").ToList()) : msg.Speech; - msg.Speech = ReplaceParameters4Response(intentResponse.Parameters, msg.Speech); + msg.Speech = ReplaceParameters4Response(intentResponse.Parameters, msg.Speech); + } } }); } @@ -285,10 +297,8 @@ namespace BotSharp.Core.Engines return client.Execute(rest); } - public static AIResponse TextRequestPerContexts(this RasaAi rasa, AIRequest request) + private static string GetModelPerContexts(RasaAi rasa, AIRequest request) { - AIResponse aiResponse = new AIResponse(); - RasaResponse response = null; Database dc = rasa.dc; // Merge input contexts @@ -316,138 +326,9 @@ namespace BotSharp.Core.Engines }).OrderByDescending(x => x.Contexts.Count).ToList(); // query per request contexts - { - string contextId = $"{String.Join(',', contexts.Select(x => x.Name))}".GetMd5Hash(); - string modelName = dc.Table().FirstOrDefault(x => x.ContextId == contextId)?.ModelName; - var result = CallRasa(rasa.agent.Id, request.Query.First(), modelName); + var contextHashs = intents.Select(x => x.ContextHash).Distinct().ToList(); - if (result.Data.Intent != null) - { - response = result.Data; - } - } - - // Max contexts match - if (response == null) - { - foreach (var it in intents) - { - request.Contexts = it.Contexts.Select(x => new AIContext { Name = x.Name.ToLower() }) - .OrderBy(x => x.Name) - .ToList(); - string contextId = $"{String.Join(',', request.Contexts.Select(x => x.Name))}".GetMd5Hash(); - - string modelName = dc.Table().FirstOrDefault(x => x.ContextId == contextId)?.ModelName; - - var result = CallRasa(rasa.agent.Id, request.Query.First(), modelName); - - if (result.Data.Intent != null) - { - response = result.Data; - break; - } - }; - } - - var intent = (dc.Table().Where(x => x.Name == response.Intent.Name) - .Include(x => x.Responses).ThenInclude(x => x.Contexts) - .Include(x => x.Responses).ThenInclude(x => x.Parameters) - .Include(x => x.Responses).ThenInclude(x => x.Messages)).First(); - - var intentResponse = ArrayHelper.GetRandom(intent.Responses); - aiResponse.Id = Guid.NewGuid().ToString(); - aiResponse.Lang = rasa.agent.Language; - aiResponse.Status = new AIResponseStatus { }; - aiResponse.SessionId = rasa.AiConfig.SessionId; - aiResponse.Timestamp = DateTime.UtcNow; - intentResponse.Messages = intentResponse.Messages.OrderBy(x => x.UpdatedTime).ToList(); - intentResponse.Messages.ToList() - .ForEach(msg => - { - if (msg.Type == AIResponseMessageType.Custom) - { - - } - else - { - msg.Speech = msg.Speech.StartsWith("[") ? - ArrayHelper.GetRandom(msg.Speech.Substring(2, msg.Speech.Length - 4).Split("\",\"").ToList()) : - msg.Speech; - } - }); - - aiResponse.Result = new AIResponseResult - { - Source = "agent", - ResolvedQuery = request.Query.First(), - Action = intentResponse.Action, - Parameters = new Dictionary(), - Score = response.Intent.Confidence, - Metadata = new AIResponseMetadata { IntentId = intent.Id, IntentName = intent.Name }, - Fulfillment = new AIResponseFulfillment - { - Messages = intentResponse.Messages.Select(x => { - if (x.Type == AIResponseMessageType.Custom) - { - return (new - { - x.Type, - x.Payload - }) as Object; - } - else - { - return (new { x.Type, x.Speech }) as Object; - } - - }).ToList() - } - }; - - // Merge context lifespan - // override if exists, otherwise add, delete if lifespan is zero - dc.DbTran(() => - { - var sessionContexts = dc.Table().Where(x => x.ConversationId == rasa.AiConfig.SessionId).ToList(); - - // minus 1 round - sessionContexts.Where(x => !intentResponse.Contexts.Select(ctx => ctx.Name).Contains(x.Context)) - .ToList() - .ForEach(ctx => ctx.Lifespan = ctx.Lifespan - 1); - - intentResponse.Contexts.ForEach(ctx => - { - var session1 = sessionContexts.FirstOrDefault(x => x.Context == ctx.Name); - - if (session1 != null) - { - if (ctx.Lifespan == 0) - { - dc.Table().Remove(session1); - } - else - { - session1.Lifespan = ctx.Lifespan; - } - } - else - { - dc.Table().Add(new ConversationContext - { - ConversationId = rasa.AiConfig.SessionId, - Context = ctx.Name, - Lifespan = ctx.Lifespan - }); - } - }); - }); - - aiResponse.Result.Contexts = dc.Table() - .Where(x => x.ConversationId == rasa.AiConfig.SessionId) - .Select(x => new AIContext { Name = x.Context.ToLower(), Lifespan = x.Lifespan }) - .ToArray(); - - return aiResponse; + return contextHashs.FirstOrDefault(); } } } diff --git a/BotSharp.Core/Intents/ContextModelMapping.cs b/BotSharp.Core/Intents/ContextModelMapping.cs deleted file mode 100644 index 0d7cb37e..00000000 --- a/BotSharp.Core/Intents/ContextModelMapping.cs +++ /dev/null @@ -1,25 +0,0 @@ -using EntityFrameworkCore.BootKit; -using System; -using System.Collections.Generic; -using System.ComponentModel.DataAnnotations; -using System.ComponentModel.DataAnnotations.Schema; -using System.Text; - -namespace BotSharp.Core.Intents -{ - [Table("Bot_ContextModelMapping")] - public class ContextModelMapping : DbRecord, IDbRecord - { - [Required] - [StringLength(36)] - public String AgentId { get; set; } - - [Required] - [StringLength(32)] - public string ContextId { get; set; } - - [Required] - [StringLength(21)] - public string ModelName { get; set; } - } -} diff --git a/BotSharp.Core/Intents/Intent.cs b/BotSharp.Core/Intents/Intent.cs index 25ee198c..162caa5d 100644 --- a/BotSharp.Core/Intents/Intent.cs +++ b/BotSharp.Core/Intents/Intent.cs @@ -36,7 +36,9 @@ namespace BotSharp.Core.Intents { get { - return $"{String.Join(',', Contexts.OrderBy(x => x.Name).Select(x => x.Name))}".GetMd5Hash(); + return Contexts == null || Contexts.Count == 0 + ? Guid.Empty.ToString("N") + : $"{String.Join(',', Contexts.OrderBy(x => x.Name).Select(x => x.Name))}".GetMd5Hash(); } } diff --git a/BotSharp.UnitTest/AgentTest.cs b/BotSharp.UnitTest/AgentTest.cs index 7c616554..766f0723 100644 --- a/BotSharp.UnitTest/AgentTest.cs +++ b/BotSharp.UnitTest/AgentTest.cs @@ -82,9 +82,7 @@ namespace BotSharp.UnitTest var rasa = new RasaAi(dc, config); - string msg = rasa.TrainWithContexts(); - - Assert.IsTrue(!String.IsNullOrEmpty(msg)); + rasa.TrainWithContexts(); } } } diff --git a/BotSharp.UnitTest/BotSharp.UnitTest.csproj b/BotSharp.UnitTest/BotSharp.UnitTest.csproj index 57854ba7..aaa14a74 100644 --- a/BotSharp.UnitTest/BotSharp.UnitTest.csproj +++ b/BotSharp.UnitTest/BotSharp.UnitTest.csproj @@ -7,11 +7,8 @@ - - - diff --git a/BotSharp.UnitTest/IntentTest.cs b/BotSharp.UnitTest/IntentTest.cs index d2c2bc83..436ccdd4 100644 --- a/BotSharp.UnitTest/IntentTest.cs +++ b/BotSharp.UnitTest/IntentTest.cs @@ -22,6 +22,29 @@ namespace BotSharp.UnitTest // Round 1 var response = rasa.TextRequest(new AIRequest { Query = new String[] { "Can you play country music?" } }); Assert.AreEqual(response.Result.Metadata.IntentName, "music.play"); + Assert.AreEqual(response.Result.Contexts.First(x => x.Name == "music-player-control").Lifespan, 3); + Assert.AreEqual(response.Result.Contexts.First(x => x.Name == "play-music").Lifespan, 5); + Assert.AreEqual(response.Result.Parameters.First(x => x.Key == "genre").Value, "country"); + + // Round 2 + response = rasa.TextRequest(new AIRequest { Query = new String[] { "pause it" } }); + Assert.AreEqual(response.Result.Metadata.IntentName, "music_player_control.pause"); + Assert.AreEqual(response.Result.Contexts.First(x => x.Name == "music-player-control").Lifespan, 3); + Assert.AreEqual(response.Result.Contexts.First(x => x.Name == "play-music").Lifespan, 4); + + // Round 3 + response = rasa.TextRequest(new AIRequest { Query = new String[] { "continue" } }); + Assert.AreEqual(response.Result.Metadata.IntentName, "music_player_control.resume"); + Assert.AreEqual(response.Result.Contexts.First(x => x.Name == "music-player-control").Lifespan, 3); + Assert.AreEqual(response.Result.Contexts.First(x => x.Name == "play-music").Lifespan, 3); + + // Round 4 + response = rasa.TextRequest(new AIRequest { Query = new String[] { "play Hard Times by David Newman" } }); + Assert.AreEqual(response.Result.Metadata.IntentName, "music.play"); + Assert.AreEqual(response.Result.Contexts.First(x => x.Name == "music-player-control").Lifespan, 3); + Assert.AreEqual(response.Result.Contexts.First(x => x.Name == "play-music").Lifespan, 5); + Assert.AreEqual(response.Result.Parameters.First(x => x.Key == "song").Value, "Hard Times"); + Assert.AreEqual(response.Result.Parameters.First(x => x.Key == "artist").Value, "David Newman"); } } } diff --git a/BotSharp.UnitTest/Settings/settings.bot.json b/BotSharp.UnitTest/Settings/settings.bot.json index 1e115890..596c6c28 100644 --- a/BotSharp.UnitTest/Settings/settings.bot.json +++ b/BotSharp.UnitTest/Settings/settings.bot.json @@ -1,6 +1,6 @@ { "Rasa": { - "Nlu": "http://localhost:5000" + "Nlu": "http://gtx.local:5000" }, "BotSharpAi": {