From ee27d804c2fb647f52971793476e1b37ef6c8e33 Mon Sep 17 00:00:00 2001 From: Bolo Peng Date: Wed, 5 Sep 2018 17:10:47 -0500 Subject: [PATCH] try doc2vec using spacy, the svm classification based on that does NOThave a good performance --- .../Engines/BotSharp/BotSharpSVMClassifier.cs | 50 ++++++++++++++++++- BotSharp.NLP/Classify/SVMClassifier.cs | 6 +-- BotSharp.WebHost/Settings/bot.json | 4 +- 3 files changed, 54 insertions(+), 6 deletions(-) diff --git a/BotSharp.Core/Engines/BotSharp/BotSharpSVMClassifier.cs b/BotSharp.Core/Engines/BotSharp/BotSharpSVMClassifier.cs index 7e97210e..d0dfcd34 100644 --- a/BotSharp.Core/Engines/BotSharp/BotSharpSVMClassifier.cs +++ b/BotSharp.Core/Engines/BotSharp/BotSharpSVMClassifier.cs @@ -3,7 +3,9 @@ using BotSharp.Core.Agents; using BotSharp.NLP.Classify; using DotNetToolkit; using Microsoft.Extensions.Configuration; +using Newtonsoft.Json; using Newtonsoft.Json.Linq; +using RestSharp; using System; using System.Collections.Generic; using System.Diagnostics; @@ -31,7 +33,20 @@ namespace BotSharp.Core.Engines.BotSharp Args args = new Args(); args.ModelFile = Path.Combine(Configuration.GetValue("BotSharpSVMClassifier:wordvec"), "wordvec_enu.bin"); LabeledFeatureSet featureSet = svmClassifier.FeatureSetsGenerator(new VectorGenerator(args).SingleSentence2Vec(doc.Sentences[0].Text), ""); + /* + // + var client = new RestClient("http://10.2.21.200:5005"); + var request = new RestRequest("doc2vec", Method.GET); + request.AddParameter("text", doc.Sentences[0].Text); + var response = client.Execute(request); + PredResult pred = JsonConvert.DeserializeObject(response.Content); + Vec vec = new Vec(); + vec.VecNodes = pred.Doc2Vec; + + LabeledFeatureSet featureSet = svmClassifier.FeatureSetsGenerator(vec, ""); + // + */ ClassifyOptions classifyOptions = new ClassifyOptions(); classifyOptions.Model = SVM.BotSharp.MachineLearning.Model.Read(Path.Combine(Settings.ModelDir, "svm_classifier_model")); classifyOptions.Transform = SVM.BotSharp.MachineLearning.RangeTransform.Read(Path.Combine(Settings.ModelDir, "transform_obj_data")); @@ -50,7 +65,6 @@ namespace BotSharp.Core.Engines.BotSharp } } - File.Delete(predictFileName); doc.Sentences[0].Intent = new TextClassificationResult @@ -86,6 +100,31 @@ namespace BotSharp.Core.Engines.BotSharp Args args = new Args(); args.ModelFile = Path.Combine(Configuration.GetValue("BotSharpSVMClassifier:wordvec"), "wordvec_enu.bin"); List featureSetList = svmClassifier.FeatureSetsGenerator(new VectorGenerator(args).Sentence2Vec(sentences), labels); + + /* + // try using spacy doc2vec + var client = new RestClient("http://10.2.21.200:5005"); + var request = new RestRequest("batchdoc2vec", Method.POST); + request.RequestFormat = DataFormat.Json; + + request.AddParameter("application/json", JsonConvert.SerializeObject(new {Sentences = sentences}), ParameterType.RequestBody); + + var response = client.Execute(request); + Result res = JsonConvert.DeserializeObject(response.Content); + + List vecs = new List(); + foreach (List cur in res.Doc2vecList) + { + Vec vec = new Vec(); + vec.VecNodes = cur; + vecs.Add(vec); + } + List featureSetList = svmClassifier.FeatureSetsGenerator(vecs, labels); + // + */ + + + ClassifyOptions classifyOptions = new ClassifyOptions(); classifyOptions.ModelFilePath = Path.Combine(Settings.ModelDir, "svm_classifier_model"); classifyOptions.TransformFilePath = Path.Combine(Settings.ModelDir, "transform_obj_data"); @@ -96,4 +135,13 @@ namespace BotSharp.Core.Engines.BotSharp return true; } } + public class Result + { + public List> Doc2vecList { get; set; } + } + + public class PredResult + { + public List Doc2Vec{ get; set; } + } } diff --git a/BotSharp.NLP/Classify/SVMClassifier.cs b/BotSharp.NLP/Classify/SVMClassifier.cs index 3cc61996..433b1c83 100644 --- a/BotSharp.NLP/Classify/SVMClassifier.cs +++ b/BotSharp.NLP/Classify/SVMClassifier.cs @@ -44,7 +44,7 @@ namespace BotSharp.NLP.Classify predict.X = GetData(featureSets).ToArray(); predict.Y = new double[1]; predict.Count = predict.X.Count(); - predict.MaxIndex = 200; + predict.MaxIndex = 300; RangeTransform transform = options.Transform; Problem scaled = transform.Scale(predict); @@ -58,13 +58,13 @@ namespace BotSharp.NLP.Classify } public void SVMClassifierTrain(List featureSets, ClassifyOptions options, SvmType svm = SvmType.C_SVC, KernelType kernel = KernelType.RBF, bool probability = true, string outputFile = null) - { + { // copy test multiclass Model Problem train = new Problem(); train.X = GetData(featureSets).ToArray(); train.Y = GetLabels(featureSets).ToArray(); train.Count = train.X.Count(); - train.MaxIndex = 200;//int.MaxValue; + train.MaxIndex = 300;//int.MaxValue; Parameter param = new Parameter(); RangeTransform transform = RangeTransform.Compute(train); diff --git a/BotSharp.WebHost/Settings/bot.json b/BotSharp.WebHost/Settings/bot.json index 30701332..013beac1 100644 --- a/BotSharp.WebHost/Settings/bot.json +++ b/BotSharp.WebHost/Settings/bot.json @@ -7,8 +7,8 @@ }, "Pipe": { - "train": "BotSharpTokenizer, BotSharpTagger, CRFsuiteEntityRecognizer, BotSharpCBOWClassifier", - "predict": "BotSharpTokenizer, BotSharpTagger, CRFsuiteEntityRecognizer, BotSharpCBOWClassifier" + "train": "BotSharpTokenizer, BotSharpTagger, CRFsuiteEntityRecognizer, BotSharpSVMClassifier", + "predict": "BotSharpTokenizer, BotSharpTagger, CRFsuiteEntityRecognizer, BotSharpSVMClassifier" }, "BotSharpSVMClassifier": {