Training data in a single config file.
This commit is contained in:
parent
10616332e7
commit
4dff8a0bf4
|
|
@ -1,5 +1,5 @@
|
|||
{
|
||||
"Rasa": {
|
||||
"Host": "http://rasa.local:5000"
|
||||
"Host": "http://gtx.local:5000"
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -30,7 +30,78 @@ namespace BotSharp.Core.Agents
|
|||
return entity.Id;
|
||||
}
|
||||
|
||||
public static RasaTrainingData GrabCorpus(this Agent agent, Database dc, List<AIContext> ctx)
|
||||
public static RasaTrainingData GrabCorpus(this Agent agent, Database dc)
|
||||
{
|
||||
var trainingData = new RasaTrainingData
|
||||
{
|
||||
Entities = new List<RasaTraningEntity>(),
|
||||
UserSays = new List<RasaIntentExpression>()
|
||||
};
|
||||
|
||||
var expressParts = new List<IntentExpressionPart>();
|
||||
|
||||
var intents = dc.Table<Intent>()
|
||||
.Include(x => x.Contexts)
|
||||
.Include(x => x.UserSays).ThenInclude(say => say.Data)
|
||||
.Where(x => x.UserSays.Count > 0)
|
||||
.ToList();
|
||||
|
||||
intents.ForEach(intent =>
|
||||
{
|
||||
intent.UserSays.ForEach(exp =>
|
||||
{
|
||||
var say = new RasaIntentExpression
|
||||
{
|
||||
Intent = intent.Name,
|
||||
Text = String.Join("", exp.Data.OrderBy(x => x.UpdatedTime).Select(x => x.Text)),
|
||||
};
|
||||
|
||||
// convert entity format
|
||||
exp.Data.Where(x => !String.IsNullOrEmpty(x.Meta))
|
||||
.ToList()
|
||||
.ForEach(x =>
|
||||
{
|
||||
int start = say.Text.IndexOf(x.Text);
|
||||
|
||||
var part = new RasaIntentExpressionPart
|
||||
{
|
||||
Value = x.Text,
|
||||
Entity = x.Alias,
|
||||
Start = start,
|
||||
End = start + x.Text.Length
|
||||
};
|
||||
|
||||
if (say.Entities == null) say.Entities = new List<RasaIntentExpressionPart>();
|
||||
say.Entities.Add(part);
|
||||
|
||||
// assemble entity synonmus
|
||||
if (!trainingData.Entities.Any(y => y.EntityType == x.Alias && y.EntityValue == x.Text))
|
||||
{
|
||||
var allSynonyms = (from e in dc.Table<Entity>()
|
||||
join ee in dc.Table<EntityEntry>() on e.Id equals ee.EntityId
|
||||
join ees in dc.Table<EntityEntrySynonym>() on ee.Id equals ees.EntityEntryId
|
||||
where e.Name == x.Alias && ee.Value == x.Text & ees.Synonym != x.Text
|
||||
select ees.Synonym).ToList();
|
||||
|
||||
var te = new RasaTraningEntity
|
||||
{
|
||||
EntityType = x.Alias,
|
||||
EntityValue = x.Text,
|
||||
Synonyms = allSynonyms
|
||||
};
|
||||
|
||||
trainingData.Entities.Add(te);
|
||||
}
|
||||
});
|
||||
|
||||
trainingData.UserSays.Add(say);
|
||||
});
|
||||
});
|
||||
|
||||
return trainingData;
|
||||
}
|
||||
|
||||
public static RasaTrainingData GrabCorpusPerContexts(this Agent agent, Database dc, List<AIContext> ctx)
|
||||
{
|
||||
var trainingData = new RasaTrainingData
|
||||
{
|
||||
|
|
|
|||
|
|
@ -54,6 +54,180 @@ namespace BotSharp.Core.Engines
|
|||
}
|
||||
}).OrderByDescending(x => x.Contexts.Count).ToList();
|
||||
|
||||
var result = CallRasa(rasa.agent.Id, request.Query.First(), rasa.agent.Id);
|
||||
|
||||
result.Data.IntentRanking = result.Data.IntentRanking.Where(x => intents.Select(i => i.Name).Contains(x.Name)).ToList();
|
||||
result.Data.Intent = result.Data.IntentRanking.First();
|
||||
response = result.Data;
|
||||
|
||||
var intent = (dc.Table<Intent>().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.Parameters.ForEach(p => {
|
||||
string query = request.Query.First();
|
||||
var entity = response.Entities.FirstOrDefault(x => x.Entity == p.Name);
|
||||
if(entity != null)
|
||||
{
|
||||
p.Value = query.Substring(entity.Start, entity.End - entity.Start);
|
||||
}
|
||||
|
||||
// fixed entity per request
|
||||
if(request.Entities != null)
|
||||
{
|
||||
var fixedEntity = request.Entities.FirstOrDefault(x => x.Name == p.Name);
|
||||
if (fixedEntity != null)
|
||||
{
|
||||
if (query.ToLower().Contains(fixedEntity.Entries.First().Value.ToLower()))
|
||||
{
|
||||
p.Value = fixedEntity.Entries.First().Value;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
});
|
||||
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 = intentResponse.Parameters.ToDictionary(x => x.Name, x=> x.Value),
|
||||
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,
|
||||
Payload = JObject.Parse(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<SessionContext>().Where(x => x.SessionId == 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<SessionContext>().Remove(session1);
|
||||
}
|
||||
else
|
||||
{
|
||||
session1.Lifespan = ctx.Lifespan;
|
||||
}
|
||||
}
|
||||
else
|
||||
{
|
||||
dc.Table<SessionContext>().Add(new SessionContext
|
||||
{
|
||||
SessionId = rasa.AiConfig.SessionId,
|
||||
Context = ctx.Name,
|
||||
Lifespan = ctx.Lifespan
|
||||
});
|
||||
}
|
||||
});
|
||||
});
|
||||
|
||||
aiResponse.Result.Contexts = dc.Table<SessionContext>()
|
||||
.Where(x => x.SessionId == rasa.AiConfig.SessionId)
|
||||
.Select(x => new AIContext { Name = x.Context.ToLower(), Lifespan = x.Lifespan })
|
||||
.ToArray();
|
||||
|
||||
return aiResponse;
|
||||
}
|
||||
|
||||
private static IRestResponse<RasaResponse> CallRasa(string projectId, string text, string model)
|
||||
{
|
||||
var client = new RestClient($"{Database.Configuration.GetSection("Rasa:Host").Value}");
|
||||
|
||||
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);
|
||||
}
|
||||
|
||||
public static AIResponse TextRequestPerContexts(this RasaAi rasa, AIRequest request)
|
||||
{
|
||||
AIResponse aiResponse = new AIResponse();
|
||||
RasaResponse response = null;
|
||||
Database dc = rasa.dc;
|
||||
|
||||
// Merge input contexts
|
||||
var contexts = dc.Table<SessionContext>()
|
||||
.Where(x => x.SessionId == rasa.AiConfig.SessionId && x.Lifespan > 0)
|
||||
.ToList()
|
||||
.Select(x => new AIContext { Name = x.Context.ToLower(), Lifespan = x.Lifespan })
|
||||
.ToList();
|
||||
|
||||
contexts.AddRange(request.Contexts.Select(x => new AIContext { Name = x.Name.ToLower(), Lifespan = x.Lifespan }));
|
||||
contexts = contexts.OrderBy(x => x.Name).ToList();
|
||||
|
||||
// search all potential intents which input context included in contexts
|
||||
var intents = rasa.agent.Intents.Where(it =>
|
||||
{
|
||||
if (contexts.Count == 0)
|
||||
{
|
||||
return it.Contexts.Count() == 0;
|
||||
}
|
||||
else
|
||||
{
|
||||
return it.Contexts.Count() > 0 &&
|
||||
it.Contexts.Count(x => contexts.Select(ctx => ctx.Name).Contains(x.Name.ToLower())) == it.Contexts.Count;
|
||||
}
|
||||
}).OrderByDescending(x => x.Contexts.Count).ToList();
|
||||
|
||||
// training per request contexts
|
||||
{
|
||||
string contextId = $"{String.Join(',', contexts.Select(x => x.Name))}".GetMd5Hash();
|
||||
|
|
@ -142,7 +316,7 @@ namespace BotSharp.Core.Engines
|
|||
Source = "agent",
|
||||
ResolvedQuery = request.Query.First(),
|
||||
Action = intentResponse.Action,
|
||||
Parameters = new Dictionary<string, object>(),
|
||||
Parameters = new Dictionary<string, string>(),
|
||||
Score = response.Intent.Confidence,
|
||||
Metadata = new AIResponseMetadata { IntentId = intent.Id, IntentName = intent.Name },
|
||||
Fulfillment = new AIResponseFulfillment
|
||||
|
|
@ -211,20 +385,64 @@ namespace BotSharp.Core.Engines
|
|||
return aiResponse;
|
||||
}
|
||||
|
||||
private static IRestResponse<RasaResponse> CallRasa(string projectId, string text, string model)
|
||||
public static string Train(this RasaAi console, Database dc)
|
||||
{
|
||||
var client = new RestClient($"{Database.Configuration.GetSection("Rasa:Host").Value}");
|
||||
var corpus = console.agent.GrabCorpus(dc);
|
||||
|
||||
var rest = new RestRequest("parse", Method.POST);
|
||||
string json = JsonConvert.SerializeObject(new { Project = projectId, Q = text, Model = model },
|
||||
// Add some fake data
|
||||
if(corpus.UserSays.Count < 3)
|
||||
{
|
||||
corpus.UserSays.Add(new RasaIntentExpression
|
||||
{
|
||||
Intent = "Welcome",
|
||||
Text = "Hi"
|
||||
});
|
||||
|
||||
corpus.UserSays.Add(new RasaIntentExpression
|
||||
{
|
||||
Intent = "Welcome",
|
||||
Text = "Hey"
|
||||
});
|
||||
|
||||
corpus.UserSays.Add(new RasaIntentExpression
|
||||
{
|
||||
Intent = "Welcome",
|
||||
Text = "Hello"
|
||||
});
|
||||
}
|
||||
|
||||
string json = JsonConvert.SerializeObject(new { rasa_nlu_data = corpus },
|
||||
new JsonSerializerSettings
|
||||
{
|
||||
ContractResolver = new CamelCasePropertyNamesContractResolver()
|
||||
ContractResolver = new CamelCasePropertyNamesContractResolver(),
|
||||
NullValueHandling = NullValueHandling.Ignore
|
||||
});
|
||||
|
||||
var client = new RestClient($"{Database.Configuration.GetSection("Rasa:Host").Value}");
|
||||
var rest = new RestRequest("train", Method.POST);
|
||||
rest.AddQueryParameter("project", console.agent.Id);
|
||||
rest.AddParameter("application/json", json, ParameterType.RequestBody);
|
||||
|
||||
return client.Execute<RasaResponse>(rest);
|
||||
var response = client.Execute(rest);
|
||||
|
||||
if (response.IsSuccessful)
|
||||
{
|
||||
var result = JObject.Parse(response.Content);
|
||||
|
||||
string modelName = result["info"].Value<String>().Split(": ")[1];
|
||||
|
||||
return modelName;
|
||||
}
|
||||
else
|
||||
{
|
||||
var result = JObject.Parse(response.Content);
|
||||
|
||||
Console.WriteLine(result["error"]);
|
||||
|
||||
return String.Empty;
|
||||
}
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Need two categories at least
|
||||
/// </summary>
|
||||
|
|
@ -235,7 +453,7 @@ namespace BotSharp.Core.Engines
|
|||
/// <returns></returns>
|
||||
public static string TrainWithContexts(this RasaAi console, Database dc, AIRequest request, String contextId)
|
||||
{
|
||||
var corpus = console.agent.GrabCorpus(dc, request.Contexts);
|
||||
var corpus = console.agent.GrabCorpusPerContexts(dc, request.Contexts);
|
||||
|
||||
corpus.UserSays.Add(new RasaIntentExpression
|
||||
{
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ namespace BotSharp.Core.Models
|
|||
}
|
||||
}
|
||||
|
||||
public Dictionary<string, object> Parameters { get; set; }
|
||||
public Dictionary<string, string> Parameters { get; set; }
|
||||
|
||||
public AIContext[] Contexts { get; set; }
|
||||
|
||||
|
|
@ -125,10 +125,10 @@ namespace BotSharp.Core.Models
|
|||
|
||||
if (Parameters.ContainsKey(name))
|
||||
{
|
||||
var parameter = Parameters[name] as JObject;
|
||||
var parameter = Parameters[name].ToString();
|
||||
if (parameter != null)
|
||||
{
|
||||
return parameter;
|
||||
return JObject.FromObject(parameter);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ namespace BotSharp.UnitTest
|
|||
public class AgentTest : TestEssential
|
||||
{
|
||||
[TestMethod]
|
||||
public void CreateAgent()
|
||||
public void CreateAgentTes()
|
||||
{
|
||||
var agent = new Agent
|
||||
{
|
||||
|
|
@ -28,7 +28,7 @@ namespace BotSharp.UnitTest
|
|||
}
|
||||
|
||||
[TestMethod]
|
||||
public void UpdateAgent()
|
||||
public void UpdateAgentTest()
|
||||
{
|
||||
var agent = new Agent
|
||||
{
|
||||
|
|
@ -42,7 +42,7 @@ namespace BotSharp.UnitTest
|
|||
}
|
||||
|
||||
[TestMethod]
|
||||
public void RestoreAgent()
|
||||
public void RestoreAgentTest()
|
||||
{
|
||||
var rasa = new RasaAi(dc);
|
||||
var importer = new AgentImporterInDialogflow();
|
||||
|
|
@ -56,9 +56,14 @@ namespace BotSharp.UnitTest
|
|||
}
|
||||
|
||||
[TestMethod]
|
||||
public void Train()
|
||||
public void TrainAgentTest()
|
||||
{
|
||||
var rasa = new RasaAi(dc);
|
||||
var config = new AIConfiguration(BOT_CLIENT_TOKEN, SupportedLanguage.English);
|
||||
config.SessionId = Guid.NewGuid().ToString();
|
||||
|
||||
var rasa = new RasaAi(dc, config);
|
||||
rasa.agent = rasa.LoadAgent();
|
||||
rasa.Train(dc);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9,9 +9,9 @@
|
|||
<ItemGroup>
|
||||
<PackageReference Include="Microsoft.Extensions.Configuration" Version="2.0.1" />
|
||||
<PackageReference Include="Microsoft.Extensions.Configuration.Json" Version="2.0.1" />
|
||||
<PackageReference Include="Microsoft.NET.Test.Sdk" Version="15.6.2" />
|
||||
<PackageReference Include="MSTest.TestAdapter" Version="1.2.0" />
|
||||
<PackageReference Include="MSTest.TestFramework" Version="1.2.0" />
|
||||
<PackageReference Include="Microsoft.NET.Test.Sdk" Version="15.7.0" />
|
||||
<PackageReference Include="MSTest.TestAdapter" Version="1.2.1" />
|
||||
<PackageReference Include="MSTest.TestFramework" Version="1.2.1" />
|
||||
</ItemGroup>
|
||||
|
||||
<ItemGroup>
|
||||
|
|
|
|||
|
|
@ -18,11 +18,8 @@ namespace BotSharp.UnitTest
|
|||
|
||||
var rasa = new RasaAi(dc, config);
|
||||
|
||||
var response = rasa.TextRequest(new AIRequest { Query = new String[] { "Create a work order for PetSmart" } });
|
||||
Assert.IsTrue(response.Result.Metadata.IntentName == "Create Work Order");
|
||||
|
||||
response = rasa.TextRequest(new AIRequest { Query = new String[] { "1010" } });
|
||||
Assert.IsTrue(response.Result.Metadata.IntentName == "Telling Store Number");
|
||||
var response = rasa.TextRequest(new AIRequest { Query = new String[] { "Hello" } });
|
||||
Assert.IsTrue(response.Result.Metadata.IntentName == "greeting");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue