BotSharp/BotSharp.NLP.UnitTest/NaiveBayesClassifierTest.cs

120 lines
3.8 KiB
C#
Raw Normal View History

using BotSharp.NLP.Classify;
using BotSharp.NLP.Corpus;
using BotSharp.NLP.Tokenize;
using Microsoft.Extensions.Configuration;
using Microsoft.VisualStudio.TestTools.UnitTesting;
using System;
using System.Collections.Generic;
using System.IO;
2018-09-09 01:47:53 +00:00
using System.Linq;
using System.Text;
2018-09-09 01:47:53 +00:00
using BotSharp.Algorithm.Extensions;
namespace BotSharp.NLP.UnitTest
{
[TestClass]
public class NaiveBayesClassifierTest : TestEssential
{
2018-09-09 03:59:01 +00:00
[TestMethod]
public void CookingTest()
{
var reader = new FasttextDataReader();
var sentences = reader.Read(new ReaderOptions
{
DataDir = Path.Combine(Configuration.GetValue<String>("MachineLearning:dataDir"), "Text Classification", "cooking.stackexchange"),
FileName = "cooking.stackexchange.txt"
});
var tokenizer = new TokenizerFactory<TreebankTokenizer>(new TokenizationOptions { }, SupportedLanguage.English);
sentences.ForEach(x => x.Words = tokenizer.Tokenize(x.Text));
sentences.Shuffle();
var options = new ClassifyOptions
{
TrainingCorpusDir = Path.Combine(Configuration.GetValue<String>("MachineLearning:dataDir"), "Text Classification", "cooking.stackexchange")
};
var classifier = new ClassifierFactory<NaiveBayesClassifier>(options, SupportedLanguage.English);
var dataset = sentences.Split(0.7M);
classifier.Train(dataset.Item1);
int correct = 0;
dataset.Item2.ForEach(td =>
{
var classes = classifier.Classify(td);
if (td.Label == classes[0].Item1)
{
correct++;
}
});
var accuracy = (float)correct / dataset.Item2.Count;
}
[TestMethod]
public void GenderTest()
{
var options = new ClassifyOptions
{
2018-09-06 22:32:51 +00:00
TrainingCorpusDir = Path.Combine(Configuration.GetValue<String>("MachineLearning:dataDir"), "Gender")
};
var classifier = new ClassifierFactory<NaiveBayesClassifier>(options, SupportedLanguage.English);
var corpus = GetLabeledCorpus(options);
var tokenizer = new TokenizerFactory<RegexTokenizer>(new TokenizationOptions
{
Pattern = RegexTokenizer.WORD_PUNC
}, SupportedLanguage.English);
corpus.ForEach(x => x.Words = tokenizer.Tokenize(x.Text));
2018-09-09 14:36:22 +00:00
classifier.Train(corpus);
string text = "Bridget";
classifier.Classify(new Sentence { Text = text, Words = tokenizer.Tokenize(text) });
2018-09-09 01:47:53 +00:00
corpus.Shuffle();
var trainingData = corpus.Skip(2000).ToList();
classifier.Train(trainingData);
2018-09-09 01:47:53 +00:00
var testData = corpus.Take(2000).ToList();
int correct = 0;
testData.ForEach(td =>
{
var classes = classifier.Classify(td);
2018-09-09 03:59:01 +00:00
if(td.Label == classes[0].Item1)
2018-09-09 01:47:53 +00:00
{
correct++;
}
});
var accuracy = (float)correct / testData.Count;
}
private List<Sentence> GetLabeledCorpus(ClassifyOptions options)
{
var reader = new LabeledPerFileNameReader();
var genders = new List<Sentence>();
var female = reader.Read(new ReaderOptions
{
DataDir = options.TrainingCorpusDir,
FileName = "female.txt"
});
genders.AddRange(female);
var male = reader.Read(new ReaderOptions
{
DataDir = options.TrainingCorpusDir,
FileName = "male.txt"
});
genders.AddRange(male);
return genders;
}
}
}