Training data in a single config file.

This commit is contained in:
haiping008@gmail.com 2018-04-06 17:39:30 -05:00
parent 10616332e7
commit 4dff8a0bf4
7 changed files with 365 additions and 74 deletions

View file

@ -1,5 +1,5 @@
{
"Rasa": {
"Host": "http://rasa.local:5000"
"Host": "http://gtx.local:5000"
}
}

View file

@ -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
{

View file

@ -54,61 +54,11 @@ namespace BotSharp.Core.Engines
}
}).OrderByDescending(x => x.Contexts.Count).ToList();
// training per request contexts
{
string contextId = $"{String.Join(',', contexts.Select(x => x.Name))}".GetMd5Hash();
string modelName = dc.Table<ContextModelMapping>().FirstOrDefault(x => x.ContextId == contextId)?.ModelName;
// need training
if (String.IsNullOrEmpty(modelName))
{
request.Contexts = contexts.Select(x => new AIContext { Name = x.Name.ToLower() })
.OrderBy(x => x.Name)
.ToList();
var result = CallRasa(rasa.agent.Id, request.Query.First(), rasa.agent.Id);
dc.DbTran(() =>
{
modelName = TrainWithContexts(rasa, dc, request, contextId);
});
}
var result = CallRasa(rasa.agent.Id, request.Query.First(), modelName);
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<ContextModelMapping>().FirstOrDefault(x => x.ContextId == contextId)?.ModelName;
// need training
if (String.IsNullOrEmpty(modelName))
{
dc.DbTran(() =>
{
modelName = TrainWithContexts(rasa, dc, request, contextId);
});
}
var result = CallRasa(rasa.agent.Id, request.Query.First(), modelName);
if (result.Data.Intent != null)
{
response = result.Data;
break;
}
};
}
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)
@ -121,6 +71,28 @@ namespace BotSharp.Core.Engines
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 =>
@ -142,7 +114,7 @@ namespace BotSharp.Core.Engines
Source = "agent",
ResolvedQuery = request.Query.First(),
Action = intentResponse.Action,
Parameters = new Dictionary<string, object>(),
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
@ -225,6 +197,252 @@ namespace BotSharp.Core.Engines
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();
string modelName = dc.Table<ContextModelMapping>().FirstOrDefault(x => x.ContextId == contextId)?.ModelName;
// need training
if (String.IsNullOrEmpty(modelName))
{
request.Contexts = contexts.Select(x => new AIContext { Name = x.Name.ToLower() })
.OrderBy(x => x.Name)
.ToList();
dc.DbTran(() =>
{
modelName = TrainWithContexts(rasa, dc, request, contextId);
});
}
var result = CallRasa(rasa.agent.Id, request.Query.First(), modelName);
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<ContextModelMapping>().FirstOrDefault(x => x.ContextId == contextId)?.ModelName;
// need training
if (String.IsNullOrEmpty(modelName))
{
dc.DbTran(() =>
{
modelName = TrainWithContexts(rasa, dc, request, contextId);
});
}
var result = CallRasa(rasa.agent.Id, request.Query.First(), modelName);
if (result.Data.Intent != null)
{
response = result.Data;
break;
}
};
}
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.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<string, string>(),
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;
}
public static string Train(this RasaAi console, Database dc)
{
var corpus = console.agent.GrabCorpus(dc);
// 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(),
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);
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
{

View file

@ -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);
}
}

View file

@ -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);
}
}
}

View file

@ -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>

View file

@ -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");
}
}
}