Files
QuanTAlib/lib/errors/mse/Mse.Validation.Tests.cs
T
2026-02-28 14:14:35 -08:00

74 lines
2.3 KiB
C#

using MathNet.Numerics;
using QuanTAlib.Tests;
namespace QuanTAlib.Validation;
public sealed class MseValidationTests : IDisposable
{
private readonly ValidationTestData _data = new();
public void Dispose() => _data.Dispose();
[Fact]
public void Mse_Matches_MathNet()
{
int[] periods = { 5, 10, 20, 50, 100 };
var quotes = _data.SkenderQuotes.ToList();
double[] actual = quotes.Select(q => (double)q.Close).ToArray();
double[] predicted = quotes.Select(q => (double)q.Open).ToArray();
foreach (int period in periods)
{
var mse = new Mse(period);
for (int i = 0; i < actual.Length; i++)
{
var val = mse.Update(
new TValue(quotes[i].Date, actual[i]),
new TValue(quotes[i].Date, predicted[i]));
// Validate last 100 bars
if (i >= actual.Length - 100 && i >= period - 1)
{
var windowActual = actual[(i - period + 1)..(i + 1)];
var windowPredicted = predicted[(i - period + 1)..(i + 1)];
double expected = Distance.MSE(windowActual, windowPredicted);
Assert.Equal(expected, val.Value, 1e-9);
}
}
}
}
[Fact]
public void Mse_Batch_Matches_MathNet()
{
int[] periods = { 5, 10, 20, 50, 100 };
var quotes = _data.SkenderQuotes.ToList();
double[] actual = quotes.Select(q => (double)q.Close).ToArray();
double[] predicted = quotes.Select(q => (double)q.Open).ToArray();
foreach (int period in periods)
{
double[] output = new double[actual.Length];
Mse.Batch(actual, predicted, output, period);
// Validate last 100 bars
for (int i = actual.Length - 100; i < actual.Length; i++)
{
if (i >= period - 1)
{
var windowActual = actual[(i - period + 1)..(i + 1)];
var windowPredicted = predicted[(i - period + 1)..(i + 1)];
double expected = Distance.MSE(windowActual, windowPredicted);
Assert.Equal(expected, output[i], 1e-9);
}
}
}
}
}