BotSharp/BotSharp.Core/Engines/Rasa/RasaAi.cs

224 lines
8.4 KiB
C#
Raw Normal View History

using BotSharp.Core.Adapters.Rasa;
using BotSharp.Core.Agents;
2018-03-28 22:08:49 +00:00
using BotSharp.Core.Entities;
2018-06-18 14:15:54 +00:00
using BotSharp.Core.Intents;
2018-03-28 22:08:49 +00:00
using BotSharp.Core.Models;
2018-06-18 22:15:32 +00:00
using DotNetToolkit;
using EntityFrameworkCore.BootKit;
using Microsoft.EntityFrameworkCore;
using Microsoft.Extensions.Configuration;
using Newtonsoft.Json;
using Newtonsoft.Json.Linq;
2018-06-13 12:23:03 +00:00
using Newtonsoft.Json.Serialization;
using RestSharp;
using System;
using System.Collections.Generic;
using System.IO;
using System.Linq;
using System.Text;
2018-03-28 22:08:49 +00:00
namespace BotSharp.Core.Engines
{
2018-04-30 03:09:32 +00:00
/// <summary>
/// Rasa nlu >= 0.12
2018-04-30 03:09:32 +00:00
/// </summary>
public class RasaAi : BotEngineBase, IBotPlatform
{
2018-07-12 15:07:05 +00:00
public AIResponse TextRequest(AIRequest request)
2018-06-13 12:23:03 +00:00
{
2018-07-12 15:07:05 +00:00
AIResponse aiResponse = new AIResponse();
2018-06-13 12:23:03 +00:00
2018-07-12 15:07:05 +00:00
string model = RasaRequestExtension.GetModelPerContexts(agent, AiConfig, request, dc);
var result = CallRasa(agent.Id, request.Query.First(), model);
2018-07-12 15:07:05 +00:00
result.Content.Log();
2018-06-13 12:23:03 +00:00
2018-07-12 15:07:05 +00:00
RasaResponse response = result.Data;
aiResponse.Id = Guid.NewGuid().ToString();
aiResponse.Lang = agent.Language;
aiResponse.Status = new AIResponseStatus { };
aiResponse.SessionId = AiConfig.SessionId;
aiResponse.Timestamp = DateTime.UtcNow;
2018-06-13 12:23:03 +00:00
2018-07-12 15:07:05 +00:00
var intentResponse = RasaRequestExtension.HandleIntentPerContextIn(agent, AiConfig, request, result.Data, dc);
2018-06-13 12:23:03 +00:00
2018-07-12 15:07:05 +00:00
RasaRequestExtension.HandleParameter(agent, intentResponse, response, request);
2018-06-13 12:23:03 +00:00
2018-07-12 15:07:05 +00:00
RasaRequestExtension.HandleMessage(intentResponse);
aiResponse.Result = new AIResponseResult
2018-06-13 12:23:03 +00:00
{
2018-07-12 15:07:05 +00:00
Source = "agent",
ResolvedQuery = request.Query.First(),
Action = intentResponse?.Action,
Parameters = intentResponse?.Parameters?.ToDictionary(x => x.Name, x => (object)x.Value),
2018-07-12 15:07:05 +00:00
Score = response.Intent.Confidence,
Metadata = new AIResponseMetadata { IntentId = intentResponse?.IntentId, IntentName = intentResponse?.IntentName },
Fulfillment = new AIResponseFulfillment
{
Messages = intentResponse?.Messages?.Select(x => {
if (x.Type == AIResponseMessageType.Custom)
{
return (new
{
x.Type,
Payload = JsonConvert.DeserializeObject(x.PayloadJson)
}) as Object;
}
else
{
return (new { x.Type, x.Speech }) as Object;
}
2018-06-13 12:23:03 +00:00
2018-07-12 15:07:05 +00:00
}).ToList()
}
};
2018-06-13 12:23:03 +00:00
2018-07-12 15:07:05 +00:00
RasaRequestExtension.HandleContext(dc, AiConfig, intentResponse, aiResponse);
2018-06-13 12:23:03 +00:00
2018-07-12 15:07:05 +00:00
Console.WriteLine(JsonConvert.SerializeObject(aiResponse.Result));
2018-06-13 12:23:03 +00:00
2018-07-12 15:07:05 +00:00
return aiResponse;
}
private IRestResponse<RasaResponse> CallRasa(string projectId, string text, string model)
{
2018-08-02 21:14:46 +00:00
var config = (IConfiguration)AppDomain.CurrentDomain.GetData("Configuration");
var client = new RestClient($"{config.GetSection("Rasa:Nlu").Value}");
2018-07-12 15:07:05 +00:00
var rest = new RestRequest("parse", Method.POST);
string json = JsonConvert.SerializeObject(new { Project = projectId, Q = text, Model = model },
new JsonSerializerSettings
{
ContractResolver = new CamelCasePropertyNamesContractResolver()
});
rest.AddParameter("application/json", json, ParameterType.RequestBody);
return client.Execute<RasaResponse>(rest);
2018-06-13 12:23:03 +00:00
}
2018-06-18 14:15:54 +00:00
2018-07-12 15:07:05 +00:00
public void Train()
2018-06-18 14:15:54 +00:00
{
var trainingData = new RasaTrainingData
{
Entities = new List<RasaTraningEntity>(),
UserSays = new List<RasaIntentExpression>()
};
var corpus = GetIntentExpressions();
2018-08-02 21:14:46 +00:00
var config = (IConfiguration)AppDomain.CurrentDomain.GetData("Configuration");
var client = new RestClient($"{config.GetSection("Rasa:Nlu").Value}");
2018-06-18 14:15:54 +00:00
2018-06-18 22:15:32 +00:00
var contextHashs = corpus.UserSays
.Select(x => x.ContextHash)
.Distinct()
.ToList();
2018-06-18 14:15:54 +00:00
2018-06-18 22:15:32 +00:00
contextHashs.ForEach(ctx =>
2018-06-18 14:15:54 +00:00
{
2018-06-21 19:14:06 +00:00
var common_examples = corpus.UserSays.Where(x => x.ContextHash == ctx || x.ContextHash == Guid.Empty.ToString("N")).ToList();
// assemble entity and synonyms
var usedEntities = new List<String>();
common_examples.ForEach(x =>
{
if (x.Entities != null)
{
usedEntities.AddRange(x.Entities.Select(y => y.Entity));
}
});
usedEntities = usedEntities.Distinct().ToList();
var entity_synonyms = corpus.Entities.Where(x => usedEntities.Contains(x.EntityType)).ToList();
2018-06-18 22:15:32 +00:00
var data = new RasaTrainingData
{
Entities = entity_synonyms.Select(x => x.ToObject<RasaTraningEntity>()).ToList(),
UserSays = common_examples.Select(x => x.ToObject<RasaIntentExpression>()).ToList()
2018-06-18 22:15:32 +00:00
};
2018-06-18 14:15:54 +00:00
2018-06-18 22:15:32 +00:00
// meet minimal requirement
// at least 2 different classes
int count = data.UserSays
.Select(x => x.Intent)
.Distinct().Count();
2018-06-18 14:15:54 +00:00
2018-06-18 22:15:32 +00:00
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 =>
2018-06-18 14:15:54 +00:00
{
2018-06-18 22:15:32 +00:00
if(data.UserSays.Count(x => x.Intent == intent) < 2)
{
data.UserSays.Add(new RasaIntentExpression
{
Intent = intent,
Text = Guid.NewGuid().ToString("N")
});
}
2018-06-18 14:15:54 +00:00
});
// set empty synonym to null
data.Entities
.Where(x => x.Synonyms != null)
.ToList()
.ForEach(entity =>
{
if (entity.Synonyms.Count == 0)
{
entity.Synonyms = null;
}
});
2018-06-18 22:15:32 +00:00
string json = JsonConvert.SerializeObject(new { rasa_nlu_data = data },
new JsonSerializerSettings
{
ContractResolver = new CamelCasePropertyNamesContractResolver(),
NullValueHandling = NullValueHandling.Ignore,
2018-06-18 22:15:32 +00:00
});
2018-06-18 14:15:54 +00:00
2018-06-18 22:15:32 +00:00
var rest = new RestRequest("train", Method.POST);
rest.AddQueryParameter("project", agent.Id);
rest.AddQueryParameter("model", ctx);
2018-07-12 15:07:05 +00:00
string trainingConfig = agent.Language == "zh" ? "config_jieba_mitie_sklearn.yml" : "config_mitie_sklearn.yml";
2018-08-02 21:14:46 +00:00
var contentRootPatch = AppDomain.CurrentDomain.GetData("ContentRootPath").ToString();
string body = File.ReadAllText(Path.Join(contentRootPatch, "Settings", trainingConfig));
2018-06-18 22:15:32 +00:00
body = $"{body}\r\ndata: {json}";
rest.AddParameter("application/x-yml", body, ParameterType.RequestBody);
var response = client.Execute(rest);
if (response.IsSuccessful)
{
var result = JObject.Parse(response.Content);
string modelName = result["info"].Value<String>().Split(": ")[1];
}
else
{
var result = JObject.Parse(response.Content);
Console.WriteLine(result["error"]);
result["error"].Log();
}
});
2018-06-18 14:15:54 +00:00
}
}
}