summaryrefslogtreecommitdiff
path: root/RhSolutions.ML.Tests/RhSolutionsTests.cs
blob: e1ec8f48fb058e263e1947ab561fe0fa5308027d (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
namespace RhSolutions.ML.Tests;

public abstract class RhSolutionsTests
{
	protected static string _appPath = Path.GetDirectoryName(Environment.GetCommandLineArgs()[0]) ?? ".";
	protected static string _dataPath = Path.Combine(_appPath, "..", "..", "..", "..", "Models", "model.zip");
	protected MLContext _mlContext;
	protected PredictionEngine<Product, TypePrediction> _predEngine;

	[SetUp]
	public void Setup()
	{
		_mlContext = new MLContext(seed: 0);
		ITransformer loadedNodel = _mlContext.Model.Load(_dataPath, out var _);
		_predEngine = _mlContext.Model.CreatePredictionEngine<Product, TypePrediction>(loadedNodel);		
	}

	public void Execute(string name, string expectedGroup)
	{
		Product p = new()
		{
			Name = name
		};
		var prediction = _predEngine.Predict(p);
		Assert.That(prediction.Type, Is.EqualTo(expectedGroup));
	}	
}